From 61aabf206b6eaa6a162694ebd4387f1939cdc0ac Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Fri, 9 Oct 2026 08:06:21 +0000 Subject: [PATCH 1/5] kernels: do not reload kernel with the same ID --- kernels/src/kernels/importer.py | 11 ++++++ kernels/tests/test_importer.py | 62 +++++++++++++++++++++++++++++++++ 2 files changed, 73 insertions(+) diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index c18d930a..c8ef40ab 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -71,6 +71,17 @@ def _import_from_path( return loaded_kernel.module metadata = Metadata.read_from_file(variant_path / "metadata.json") + + # Kernel ids are unique per build: if this build was already imported + # reuse it instead of executing it again. + if (module := sys.modules.get(metadata.id)) is not None: + _loaded_kernels[variant_path] = LoadedKernel( + metadata=metadata, + module=module, + repo_info=repo_info, + ) + return module + module_name = metadata.name.python_name file_path = variant_path / "__init__.py" diff --git a/kernels/tests/test_importer.py b/kernels/tests/test_importer.py index 3acb0bbb..7fd53545 100644 --- a/kernels/tests/test_importer.py +++ b/kernels/tests/test_importer.py @@ -1,10 +1,14 @@ import json import sys +import types import pytest from kernels.importer import _import_from_path, _loaded_kernels +_EXEC_LOG_MODULE = "_kernels_test_exec_log" +_COUNTING_ID = "counting_1_cuda" + def _write_variant(tmp_path): variant_dir = tmp_path / "build" / "torch28-cxx11-cu128-x86_64-linux" @@ -35,3 +39,61 @@ def test_failed_import_cleans_up_sys_modules(tmp_path): finally: _loaded_kernels.pop(variant_dir, None) sys.modules.pop("broken_1_cuda", None) + + +def _write_counting_variant(base_path, kernel_id): + """Write a kernel variant that records each execution of its module.""" + variant_dir = base_path / "build" / "torch28-cxx11-cu128-x86_64-linux" + variant_dir.mkdir(parents=True) + metadata = { + "id": kernel_id, + "name": "counting", + "version": 1, + "license": "Apache-2.0", + "python-depends": ["torch"], + "backend": {"type": "cuda"}, + } + (variant_dir / "metadata.json").write_text(json.dumps(metadata)) + (variant_dir / "__init__.py").write_text(f"import {_EXEC_LOG_MODULE}\n{_EXEC_LOG_MODULE}.calls.append(__file__)\n") + return variant_dir + + +@pytest.fixture +def exec_log(monkeypatch): + log = types.ModuleType(_EXEC_LOG_MODULE) + log.calls = [] # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, _EXEC_LOG_MODULE, log) + return log.calls # type: ignore[attr-defined] + + +def test_same_id_different_path_is_not_reloaded(tmp_path, exec_log): + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID) + second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID) + try: + first = _import_from_path(first_dir, deps={}) + second = _import_from_path(second_dir, deps={}) + + assert first is second + assert len(exec_log) == 1 + assert _loaded_kernels[first_dir].module is first + assert _loaded_kernels[second_dir].module is first + finally: + _loaded_kernels.pop(first_dir, None) + _loaded_kernels.pop(second_dir, None) + sys.modules.pop(_COUNTING_ID, None) + + +def test_already_imported_kernel_is_reregistered(tmp_path, exec_log): + variant_dir = _write_counting_variant(tmp_path, _COUNTING_ID) + try: + first = _import_from_path(variant_dir, deps={}) + _loaded_kernels.pop(variant_dir) + + second = _import_from_path(variant_dir, deps={}) + + assert first is second + assert len(exec_log) == 1 + assert _loaded_kernels[variant_dir].module is first + finally: + _loaded_kernels.pop(variant_dir, None) + sys.modules.pop(_COUNTING_ID, None) From 9bfef44f66716c4255410b8003bccdba3f98dfe1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Fri, 9 Oct 2026 12:50:48 +0000 Subject: [PATCH 2/5] kernels: log when when a kernel was already loaded --- kernels/src/kernels/importer.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index c8ef40ab..ba32d72d 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -1,4 +1,5 @@ import importlib +import logging import sys from dataclasses import dataclass from pathlib import Path @@ -7,6 +8,8 @@ from kernels._rust import Metadata from kernels.hf_hub import RepoInfo +logger = logging.getLogger(__name__) + @dataclass(frozen=True) class LoadedKernel: @@ -75,6 +78,7 @@ def _import_from_path( # Kernel ids are unique per build: if this build was already imported # reuse it instead of executing it again. if (module := sys.modules.get(metadata.id)) is not None: + logging.debug(f"Kernel already loaded, skipping: {metadata.id}") _loaded_kernels[variant_path] = LoadedKernel( metadata=metadata, module=module, From 6fcbf6abe86b6f96d060b1d9d2c09b75e1d619ab Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Fri, 9 Oct 2026 13:56:50 +0000 Subject: [PATCH 3/5] kernels-common: split VerificationReceipt into {Signature,Digest}VerificationReceipt --- kernels-common/src/signing/receipt.rs | 199 ++++++++++++++++++++------ kernels/rust/lib.rs | 11 +- kernels/rust/signing.rs | 155 +++++++++++++++----- kernels/src/kernels/_rust.pyi | 89 +++++++++--- kernels/src/kernels/verify.py | 12 +- kernels/tests/test_verify.py | 4 +- 6 files changed, 362 insertions(+), 108 deletions(-) diff --git a/kernels-common/src/signing/receipt.rs b/kernels-common/src/signing/receipt.rs index cff9b99c..dd6e5193 100644 --- a/kernels-common/src/signing/receipt.rs +++ b/kernels-common/src/signing/receipt.rs @@ -1,7 +1,9 @@ use std::fs; use std::io::{self, Write as _}; -use std::path::PathBuf; +use std::marker::PhantomData; +use std::path::{Path, PathBuf}; +use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use thiserror::Error; @@ -9,23 +11,55 @@ use thiserror::Error; use crate::git::Oid; use crate::hf::{UnknownCacheDir, kernels_cache}; -/// Version of the on-disk receipt format. -pub const CACHE_FORMAT_VERSION: &str = "v1"; +/// Directory in the kernels cache that holds all receipt stores. +const RECEIPTS_DIR: &str = ".verified-kernels"; -/// Receipt of a successful kernel verification. +/// Version of the on-disk signature receipt format. +pub const SIGNATURE_RECEIPT_FORMAT_VERSION: &str = "v1"; + +/// Version of the on-disk digest receipt format. +pub const DIGEST_RECEIPT_FORMAT_VERSION: &str = "v1"; + +/// Receipt of a successful verification of a kernel. +pub trait Receipt: Serialize + DeserializeOwned { + /// The kernel location the verification applies to. + fn location(&self) -> &KernelLocation; +} + +/// Receipt of a successful signature verification of a kernel. #[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] -pub struct VerificationReceipt { +pub struct SignatureReceipt { /// The kernel location the receipt applies to. location: KernelLocation, } -impl VerificationReceipt { +impl SignatureReceipt { pub fn new(location: KernelLocation) -> Self { - VerificationReceipt { location } + SignatureReceipt { location } } +} - /// The kernel location the verification applies to. - pub fn location(&self) -> &KernelLocation { +impl Receipt for SignatureReceipt { + fn location(&self) -> &KernelLocation { + &self.location + } +} + +/// Receipt of a successful digest verification of a kernel. +#[derive(Clone, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)] +pub struct DigestReceipt { + /// The kernel location the receipt applies to. + location: KernelLocation, +} + +impl DigestReceipt { + pub fn new(location: KernelLocation) -> Self { + DigestReceipt { location } + } +} + +impl Receipt for DigestReceipt { + fn location(&self) -> &KernelLocation { &self.location } } @@ -74,28 +108,61 @@ impl KernelLocation { } } -/// Storage for verification receipts. +/// Storage for verification receipts of type `R`. /// /// Receipts are stored as JSON files in a cache directory, named by their /// receipt key. #[derive(Clone, Debug)] -pub struct ReceiptStore { +pub struct ReceiptStore { dir: PathBuf, + receipt: PhantomData, +} + +/// Store for signature verification receipts. +pub type SignatureReceiptStore = ReceiptStore; + +/// Store for digest verification receipts. +pub type DigestReceiptStore = ReceiptStore; + +impl SignatureReceiptStore { + /// The signature receipt store inside the kernels cache. + pub fn in_kernels_cache() -> Result { + Ok(Self::in_cache_dir(&kernels_cache()?)) + } + + fn in_cache_dir(cache_dir: &Path) -> Self { + Self::from_path( + cache_dir + .join(RECEIPTS_DIR) + .join("signature") + .join(SIGNATURE_RECEIPT_FORMAT_VERSION), + ) + } } -impl ReceiptStore { - /// The receipt store inside the kernels cache. +impl DigestReceiptStore { + /// The digest receipt store inside the kernels cache. pub fn in_kernels_cache() -> Result { - let default_dir = kernels_cache()? - .join(".verified-kernels") - .join(CACHE_FORMAT_VERSION); + Ok(Self::in_cache_dir(&kernels_cache()?)) + } - Ok(Self::from_path(default_dir)) + fn in_cache_dir(cache_dir: &Path) -> Self { + Self::from_path( + cache_dir + .join(RECEIPTS_DIR) + .join("digest") + .join(DIGEST_RECEIPT_FORMAT_VERSION), + ) } +} +impl ReceiptStore { /// A receipt store in the given directory. pub fn from_path(dir: impl Into) -> Self { - ReceiptStore { dir: dir.into() } + ReceiptStore { + dir: dir.into(), + receipt: PhantomData, + } } /// Load the receipt for the given kernel location. @@ -103,26 +170,23 @@ impl ReceiptStore { /// Returns `Ok(None)` when no receipt exists for the location, and an /// error when a receipt exists but cannot be read, is corrupt, or /// describes a different kernel. - pub fn load( - &self, - location: &KernelLocation, - ) -> Result, ReceiptStoreError> { + pub fn load(&self, location: &KernelLocation) -> Result, ReceiptStoreError> { let path = self.dir.join(location.receipt_key()); let data = match fs::read(&path) { Ok(data) => data, Err(e) if e.kind() == io::ErrorKind::NotFound => return Ok(None), Err(source) => return Err(ReceiptStoreError::Read { path, source }), }; - let receipt: VerificationReceipt = + let receipt: R = serde_json::from_slice(&data).map_err(|source| ReceiptStoreError::Corrupt { path: path.clone(), source, })?; - if receipt.location != *location { + if receipt.location() != location { return Err(ReceiptStoreError::LocationMismatch { path, - found: Box::new(receipt.location), + found: Box::new(receipt.location().clone()), }); } @@ -130,8 +194,8 @@ impl ReceiptStore { } /// Store a receipt. - pub fn store(&self, receipt: &VerificationReceipt) -> Result<(), ReceiptStoreError> { - let path = self.dir.join(receipt.location.receipt_key()); + pub fn store(&self, receipt: &R) -> Result<(), ReceiptStoreError> { + let path = self.dir.join(receipt.location().receipt_key()); let write_err = |source: io::Error| ReceiptStoreError::Write { path: path.clone(), source, @@ -214,6 +278,8 @@ fn hex_encode(bytes: &[u8]) -> String { mod tests { use std::str::FromStr; + use std::fmt::Debug; + use super::*; use tempfile::TempDir; @@ -265,11 +331,16 @@ mod tests { } #[test] - fn receipt_roundtrip() -> io::Result<()> { - let dir = TempDir::new()?; - let store = ReceiptStore::from_path(dir.path()); + fn receipt_roundtrip() { + check_receipt_roundtrip(SignatureReceipt::new); + check_receipt_roundtrip(DigestReceipt::new); + } + + fn check_receipt_roundtrip(new: fn(KernelLocation) -> R) { + let dir = TempDir::new().unwrap(); + let store = ReceiptStore::::from_path(dir.path()); let location = hub_location(); - let receipt = VerificationReceipt::new(location.clone()); + let receipt = new(location.clone()); store.store(&receipt).expect("receipt should store"); let loaded = store @@ -277,14 +348,18 @@ mod tests { .expect("receipt should load") .expect("receipt should exist"); - assert_eq!(loaded.location(), receipt.location()); - Ok(()) + assert_eq!(loaded, receipt); } #[test] fn load_receipt_missing_or_corrupt() { + check_load_receipt_missing_or_corrupt::(); + check_load_receipt_missing_or_corrupt::(); + } + + fn check_load_receipt_missing_or_corrupt() { let dir = TempDir::new().unwrap(); - let store = ReceiptStore::from_path(dir.path()); + let store = ReceiptStore::::from_path(dir.path()); let location = hub_location(); let key = location.receipt_key(); @@ -323,8 +398,15 @@ mod tests { #[test] fn load_receipt_rejects_transplanted_receipt() { + check_load_receipt_rejects_transplanted_receipt(SignatureReceipt::new); + check_load_receipt_rejects_transplanted_receipt(DigestReceipt::new); + } + + fn check_load_receipt_rejects_transplanted_receipt( + new: fn(KernelLocation) -> R, + ) { let dir = TempDir::new().unwrap(); - let store = ReceiptStore::from_path(dir.path()); + let store = ReceiptStore::::from_path(dir.path()); let signed = hub_location(); let unsigned = KernelLocation::remote( @@ -334,7 +416,7 @@ mod tests { ); store - .store(&VerificationReceipt::new(signed.clone())) + .store(&new(signed.clone())) .expect("receipt should store"); // Transplant the receipt onto the other revision's key. @@ -357,21 +439,49 @@ mod tests { #[test] fn store_receipt_fails_when_cache_not_writable() { + check_store_receipt_fails_when_cache_not_writable(SignatureReceipt::new); + check_store_receipt_fails_when_cache_not_writable(DigestReceipt::new); + } + + fn check_store_receipt_fails_when_cache_not_writable(new: fn(KernelLocation) -> R) { let dir = TempDir::new().unwrap(); let receipt_dir = dir.path().join("receipts"); // A regular file where the receipt directory should be. fs::write(&receipt_dir, "not a directory").unwrap(); - let receipt = VerificationReceipt::new(hub_location()); + let receipt = new(hub_location()); assert!(matches!( - ReceiptStore::from_path(&receipt_dir).store(&receipt), + ReceiptStore::::from_path(&receipt_dir).store(&receipt), Err(ReceiptStoreError::Write { .. }) )); } + #[test] + fn stores_use_separate_directories_in_cache() { + let cache_dir = Path::new("/cache"); + + assert_eq!( + SignatureReceiptStore::in_cache_dir(cache_dir).dir, + Path::new("/cache/.verified-kernels/signature/v1") + ); + assert_eq!( + DigestReceiptStore::in_cache_dir(cache_dir).dir, + Path::new("/cache/.verified-kernels/digest/v1") + ); + } + + /// The on-disk format must not change by accident. Changing it requires + /// bumping the format version of the store. #[test] fn receipt_json_format_is_stable() { - let receipt = VerificationReceipt::new(hub_location()); + check_receipt_json_format_is_stable(SignatureReceipt::new); + check_receipt_json_format_is_stable(DigestReceipt::new); + } + + fn check_receipt_json_format_is_stable( + new: fn(KernelLocation) -> R, + ) { + let receipt = new(hub_location()); let json = serde_json::to_string(&receipt).unwrap(); let expected = format!( @@ -381,14 +491,19 @@ mod tests { assert_eq!(json, expected); // The pinned format must roundtrip. - let parsed: VerificationReceipt = serde_json::from_str(&json).unwrap(); - assert_eq!(parsed.location(), receipt.location()); + let parsed: R = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, receipt); } #[test] fn unknown_receipt_fields_are_ignored() { + check_unknown_receipt_fields_are_ignored::(); + check_unknown_receipt_fields_are_ignored::(); + } + + fn check_unknown_receipt_fields_are_ignored() { let dir = TempDir::new().unwrap(); - let store = ReceiptStore::from_path(dir.path()); + let store = ReceiptStore::::from_path(dir.path()); let location = hub_location(); fs::write( diff --git a/kernels/rust/lib.rs b/kernels/rust/lib.rs index a502a911..6b8dd716 100644 --- a/kernels/rust/lib.rs +++ b/kernels/rust/lib.rs @@ -21,7 +21,10 @@ mod version; use config::{PyBuild, PyGeneral}; use git::PyOid; use lock::{PyKernelLock, PyKernelLocks, PyKernelPaths, PyNixKernelLock, PyNixKernelLocks}; -use signing::{PyKernelLocation, PyReceiptStore, PyVerificationReceipt, ReceiptError}; +use signing::{ + PyDigestReceipt, PyDigestReceiptStore, PyKernelLocation, PySignatureReceipt, + PySignatureReceiptStore, ReceiptError, +}; use version::PyVersion; /// A validated kernel name matching `^[a-z][-a-z0-9]*[a-z0-9]$`. @@ -778,8 +781,10 @@ 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::()?; + m.add_class::()?; + m.add_class::()?; m.add( "DigestValidationError", m.py().get_type::(), diff --git a/kernels/rust/signing.rs b/kernels/rust/signing.rs index 8eb2ebe8..7b96da89 100644 --- a/kernels/rust/signing.rs +++ b/kernels/rust/signing.rs @@ -1,6 +1,9 @@ use std::path::PathBuf; -use kernels_common::signing::receipt::{KernelLocation, ReceiptStore, VerificationReceipt}; +use kernels_common::signing::receipt::{ + DigestReceipt, DigestReceiptStore, KernelLocation, Receipt, ReceiptStoreError, + SignatureReceipt, SignatureReceiptStore, +}; use pyo3::exceptions::PyException; use pyo3::prelude::*; @@ -10,14 +13,18 @@ pyo3::create_exception!( _rust, ReceiptError, PyException, - "Raised by `ReceiptStore` when a receipt cannot be read, written, or \ - interpreted.\n\n\ - A missing receipt is not an error: `ReceiptStore.load` returns `None` \ - for it. Since a receipt is only a cache of a previous verification, \ - callers can treat this exception as a cache miss and re-verify, at the \ - cost of not noticing a cache that is persistently broken." + "Raised by `SignatureReceiptStore` and `DigestReceiptStore` when a \ + receipt cannot be read, written, or interpreted.\n\n\ + A missing receipt is not an error: `load` returns `None` for it. Since \ + a receipt is only a cache of a previous verification, callers can treat \ + this exception as a cache miss and re-verify, at the cost of not \ + noticing a cache that is persistently broken." ); +fn receipt_error(err: ReceiptStoreError) -> PyErr { + ReceiptError::new_err(format!("{:#}", eyre::Report::new(err))) +} + /// The location of a kernel that a verification applies to. #[pyclass(name = "KernelLocation", frozen, eq, hash)] #[derive(Clone, Debug, Eq, Hash, PartialEq)] @@ -55,24 +62,24 @@ impl PyKernelLocation { } } -/// Receipt of a successful kernel verification. -#[pyclass(name = "VerificationReceipt", frozen, eq, hash)] +/// Receipt of a successful signature verification of a kernel. +#[pyclass(name = "SignatureReceipt", frozen, eq, hash)] #[derive(Clone, Debug, Eq, Hash, PartialEq)] -pub(crate) struct PyVerificationReceipt { - inner: VerificationReceipt, +pub(crate) struct PySignatureReceipt { + inner: SignatureReceipt, } -impl From for PyVerificationReceipt { - fn from(inner: VerificationReceipt) -> Self { +impl From for PySignatureReceipt { + fn from(inner: SignatureReceipt) -> Self { Self { inner } } } #[pymethods] -impl PyVerificationReceipt { +impl PySignatureReceipt { #[new] fn new(location: &PyKernelLocation) -> Self { - VerificationReceipt::new(location.inner.clone()).into() + SignatureReceipt::new(location.inner.clone()).into() } #[getter] @@ -81,58 +88,130 @@ impl PyVerificationReceipt { } fn __repr__(&self) -> String { - format!( - "VerificationReceipt(location={})", - self.location().__repr__() - ) + format!("SignatureReceipt(location={})", self.location().__repr__()) } } -/// Store of kernel verification receipts. -#[pyclass(name = "ReceiptStore", frozen)] +/// Store of kernel signature verification receipts. +#[pyclass(name = "SignatureReceiptStore", frozen)] #[derive(Clone, Debug)] -pub(crate) struct PyReceiptStore { - inner: ReceiptStore, +pub(crate) struct PySignatureReceiptStore { + inner: SignatureReceiptStore, } #[pymethods] -impl PyReceiptStore { - /// The receipt store inside the kernels cache. +impl PySignatureReceiptStore { + /// The signature receipt store inside the kernels cache. /// /// Raises `ReceiptError` when the cache directory cannot be determined, /// in which case verifications cannot be cached. #[staticmethod] fn in_kernels_cache() -> PyResult { - ReceiptStore::in_kernels_cache() - .map(|inner| PyReceiptStore { inner }) - .map_err(|err| ReceiptError::new_err(format!("{:#}", eyre::Report::new(err)))) + SignatureReceiptStore::in_kernels_cache() + .map(|inner| PySignatureReceiptStore { inner }) + .map_err(receipt_error) } - /// A receipt store in the given directory. + /// A signature receipt store in the given directory. #[staticmethod] fn from_path(path: PathBuf) -> Self { - PyReceiptStore { - inner: ReceiptStore::from_path(path), + PySignatureReceiptStore { + inner: SignatureReceiptStore::from_path(path), } } - /// The receipt for `location`, or `None` when the kernel has not been - /// verified yet. + /// The receipt for `location`, or `None` when the kernel signature has + /// not been verified yet. /// /// Raises `ReceiptError` if a receipt exists but cannot be used. - fn load(&self, location: &PyKernelLocation) -> PyResult> { + fn load(&self, location: &PyKernelLocation) -> PyResult> { self.inner .load(&location.inner) .map(|receipt| receipt.map(Into::into)) - .map_err(|err| ReceiptError::new_err(format!("{:#}", eyre::Report::new(err)))) + .map_err(receipt_error) } /// Store `receipt`, replacing any existing receipt for its location. /// /// Raises `ReceiptError` if the receipt cannot be written. - fn store(&self, receipt: &PyVerificationReceipt) -> PyResult<()> { + fn store(&self, receipt: &PySignatureReceipt) -> PyResult<()> { + self.inner.store(&receipt.inner).map_err(receipt_error) + } +} + +/// Receipt of a successful digest verification of a kernel. +#[pyclass(name = "DigestReceipt", frozen, eq, hash)] +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub(crate) struct PyDigestReceipt { + inner: DigestReceipt, +} + +impl From for PyDigestReceipt { + fn from(inner: DigestReceipt) -> Self { + Self { inner } + } +} + +#[pymethods] +impl PyDigestReceipt { + #[new] + fn new(location: &PyKernelLocation) -> Self { + DigestReceipt::new(location.inner.clone()).into() + } + + #[getter] + fn location(&self) -> PyKernelLocation { + self.inner.location().clone().into() + } + + fn __repr__(&self) -> String { + format!("DigestReceipt(location={})", self.location().__repr__()) + } +} + +/// Store of kernel digest verification receipts. +#[pyclass(name = "DigestReceiptStore", frozen)] +#[derive(Clone, Debug)] +pub(crate) struct PyDigestReceiptStore { + inner: DigestReceiptStore, +} + +#[pymethods] +impl PyDigestReceiptStore { + /// The digest receipt store inside the kernels cache. + /// + /// Raises `ReceiptError` when the cache directory cannot be determined, + /// in which case verifications cannot be cached. + #[staticmethod] + fn in_kernels_cache() -> PyResult { + DigestReceiptStore::in_kernels_cache() + .map(|inner| PyDigestReceiptStore { inner }) + .map_err(receipt_error) + } + + /// A digest receipt store in the given directory. + #[staticmethod] + fn from_path(path: PathBuf) -> Self { + PyDigestReceiptStore { + inner: DigestReceiptStore::from_path(path), + } + } + + /// The receipt for `location`, or `None` when the kernel digest has not + /// been verified yet. + /// + /// Raises `ReceiptError` if a receipt exists but cannot be used. + fn load(&self, location: &PyKernelLocation) -> PyResult> { self.inner - .store(&receipt.inner) - .map_err(|err| ReceiptError::new_err(format!("{:#}", eyre::Report::new(err)))) + .load(&location.inner) + .map(|receipt| receipt.map(Into::into)) + .map_err(receipt_error) + } + + /// Store `receipt`, replacing any existing receipt for its location. + /// + /// Raises `ReceiptError` if the receipt cannot be written. + fn store(&self, receipt: &PyDigestReceipt) -> PyResult<()> { + self.inner.store(&receipt.inner).map_err(receipt_error) } } diff --git a/kernels/src/kernels/_rust.pyi b/kernels/src/kernels/_rust.pyi index 626d4206..197193e1 100644 --- a/kernels/src/kernels/_rust.pyi +++ b/kernels/src/kernels/_rust.pyi @@ -29,8 +29,10 @@ __all__ = [ "DigestViolation", "DigestValidationError", "KernelLocation", - "VerificationReceipt", - "ReceiptStore", + "SignatureReceipt", + "SignatureReceiptStore", + "DigestReceipt", + "DigestReceiptStore", "ReceiptError", "Version", "__version__", @@ -331,13 +333,13 @@ class KernelLocation: def __repr__(self) -> str: ... @final -class VerificationReceipt: - """Receipt of a successful kernel verification. +class SignatureReceipt: + """Receipt of a successful signature verification of a kernel. - A receipt states that a kernel has already been verified. If the kernel - location changed, the receipt's hash will not match anymore.""" + A receipt states that a kernel's signature has already been verified. If + the kernel location changed, the receipt's hash will not match anymore.""" - def __new__(cls, location: KernelLocation) -> "VerificationReceipt": ... + def __new__(cls, location: KernelLocation) -> "SignatureReceipt": ... @property def location(self) -> KernelLocation: """The kernel location the verification applies to.""" @@ -346,12 +348,12 @@ class VerificationReceipt: def __repr__(self) -> str: ... @final -class ReceiptStore: - """Store of kernel verification receipts.""" +class SignatureReceiptStore: + """Store of kernel signature verification receipts.""" @staticmethod - def in_kernels_cache() -> "ReceiptStore": - """The receipt store inside the kernels cache. + def in_kernels_cache() -> "SignatureReceiptStore": + """The signature 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. @@ -362,19 +364,72 @@ class ReceiptStore: ... @staticmethod - def from_path(path: os.PathLike[str] | str) -> "ReceiptStore": - """A receipt store in the given directory.""" + def from_path(path: os.PathLike[str] | str) -> "SignatureReceiptStore": + """A signature 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. + def load(self, location: KernelLocation) -> Optional[SignatureReceipt]: + """The receipt for `location`, or `None` when the kernel signature has not been verified yet. Raises: ReceiptError: If a receipt exists but cannot be used. """ ... - def store(self, receipt: VerificationReceipt) -> None: + def store(self, receipt: SignatureReceipt) -> None: + """Store `receipt`, replacing any existing receipt for its location. + + Raises: + ReceiptError: If the receipt cannot be written. + """ + ... + +@final +class DigestReceipt: + """Receipt of a successful digest verification of a kernel. + + A receipt states that a kernel's files have already been verified against + the digest in its metadata. If the kernel location changed, the receipt's + hash will not match anymore.""" + + def __new__(cls, location: KernelLocation) -> "DigestReceipt": ... + @property + def location(self) -> KernelLocation: + """The kernel location the verification applies to.""" + ... + + def __repr__(self) -> str: ... + +@final +class DigestReceiptStore: + """Store of kernel digest verification receipts.""" + + @staticmethod + def in_kernels_cache() -> "DigestReceiptStore": + """The digest 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) -> "DigestReceiptStore": + """A digest receipt store in the given directory.""" + ... + + def load(self, location: KernelLocation) -> Optional[DigestReceipt]: + """The receipt for `location`, or `None` when the kernel digest has not been verified yet. + + Raises: + ReceiptError: If a receipt exists but cannot be used. + """ + ... + + def store(self, receipt: DigestReceipt) -> None: """Store `receipt`, replacing any existing receipt for its location. Raises: @@ -383,7 +438,7 @@ class ReceiptStore: ... class ReceiptError(Exception): - """Raised by `ReceiptStore` when a receipt cannot be read, written, or interpreted.""" + """Raised by `SignatureReceiptStore` and `DigestReceiptStore` when a receipt cannot be read, written, or interpreted.""" class KernelVersion: """A kernel version: either a numeric version or a git revision string.""" diff --git a/kernels/src/kernels/verify.py b/kernels/src/kernels/verify.py index ba043e29..cd5af504 100644 --- a/kernels/src/kernels/verify.py +++ b/kernels/src/kernels/verify.py @@ -17,8 +17,8 @@ KernelLocation, Metadata, ReceiptError, - ReceiptStore, - VerificationReceipt, + SignatureReceipt, + SignatureReceiptStore, ) logger = logging.getLogger(__name__) @@ -195,16 +195,16 @@ def __str__(self) -> str: ) -def _open_receipt_store() -> ReceiptStore | None: +def _open_receipt_store() -> SignatureReceiptStore | None: """The receipt store, or `None` when verifications cannot be cached.""" try: - return ReceiptStore.in_kernels_cache() + return SignatureReceiptStore.in_kernels_cache() except ReceiptError as e: logger.warning(f"Cannot cache kernel verifications: {e}") return None -def _has_receipt(store: ReceiptStore, location: KernelLocation) -> bool: +def _has_receipt(store: SignatureReceiptStore, location: KernelLocation) -> bool: """Whether the kernel at `location` was verified before. An unusable receipt counts as a cache miss: the kernel is then verified in @@ -314,7 +314,7 @@ def verify_variant( if receipt_store is not None: try: - receipt_store.store(VerificationReceipt(location)) + receipt_store.store(SignatureReceipt(location)) except ReceiptError as e: logger.warning(f"Cannot store kernel verification receipt: {e}") diff --git a/kernels/tests/test_verify.py b/kernels/tests/test_verify.py index a942abcb..39f4a5ca 100644 --- a/kernels/tests/test_verify.py +++ b/kernels/tests/test_verify.py @@ -7,7 +7,7 @@ import kernels.verify as verify_module from kernels import install_kernel -from kernels._rust import DigestViolation, KernelLocation, Oid, ReceiptStore +from kernels._rust import DigestViolation, KernelLocation, Oid, SignatureReceiptStore from kernels._versions import resolve_revision_or_version from kernels.hf_hub import _get_cache_dir, _get_hf_api from kernels.resolver import _BYTECODE_IGNORE_PATTERNS @@ -26,7 +26,7 @@ def receipt_store(tmp_path, monkeypatch): """An isolated receipt store, so that tests do not share verifications.""" receipt_dir = tmp_path / "receipts" - store = ReceiptStore.from_path(receipt_dir) + store = SignatureReceiptStore.from_path(receipt_dir) monkeypatch.setattr(verify_module, "_open_receipt_store", lambda: store) return store From f2af595a641ab1e30d1dcae60cf35f55ae7d1d6e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Fri, 9 Oct 2026 14:36:41 +0000 Subject: [PATCH 4/5] kernels: add `DigestValidator` Also remove digest checking from `SignatureValidator`. --- docs/source/cli-verify-signature.md | 4 +- docs/source/security.md | 40 ++- kernels/src/kernels/cli/verify_signature.py | 40 ++- kernels/src/kernels/digest.py | 139 ++++++++++ kernels/src/kernels/validate.py | 58 +++- kernels/src/kernels/verify.py | 116 ++------ kernels/tests/test_digest.py | 283 ++++++++++++++++++++ kernels/tests/test_validate.py | 120 ++++++++- kernels/tests/test_verify.py | 153 ++++++----- 9 files changed, 756 insertions(+), 197 deletions(-) create mode 100644 kernels/src/kernels/digest.py create mode 100644 kernels/tests/test_digest.py diff --git a/docs/source/cli-verify-signature.md b/docs/source/cli-verify-signature.md index 396d0791..140afa7c 100644 --- a/docs/source/cli-verify-signature.md +++ b/docs/source/cli-verify-signature.md @@ -37,9 +37,9 @@ kernels verify-signature kernels-community/relu 1 --all-variants ```bash $ kernels verify-signature kernels-community/relu 1 -✅ torch211-cxx11-cu126-x86_64-linux: kernel metadata is correctly signed +✅ torch211-cxx11-cu126-x86_64-linux: the metadata is correctly signed and the files match the digest in the metadata $ kernels verify-signature kernels-community/flash-attn2 1 -❌ torch211-cxx11-cu126-x86_64-linux: cannot verify kernel integrity, signature not found +❌ torch211-cxx11-cu126-x86_64-linux: not signed, so its integrity cannot be verified ``` ## Options diff --git a/docs/source/security.md b/docs/source/security.md index 67f15598..7b546524 100644 --- a/docs/source/security.md +++ b/docs/source/security.md @@ -151,10 +151,11 @@ platform uses a Git implementation with SHA-1 collision detection `kernels` can verify kernels with cosign. -On load, `kernels` checks that the files match signed digests in -`metadata.json`. Signing uses cosign with short-lived keys, and the signature is -recorded in a [ledger](https://docs.sigstore.dev/logging/overview/). That -combination makes leaked CI signing keys much harder to reuse. +On load, `kernels` checks that the files match the digests in `metadata.json` +and that `metadata.json` is signed. Signing uses cosign with short-lived keys, +and the signature is recorded in a +[ledger](https://docs.sigstore.dev/logging/overview/). That combination makes +leaked CI signing keys much harder to reuse. The builder computes the SHA-256 digest of each file in the kernel and stores it in `metadata.json`: @@ -175,32 +176,45 @@ Aside from the main signature, cosign also records information about how the signature was made, such as the OIDC issuer, the source repository, and the workflow path/branch. -Signature verification performs the following steps: +Kernels are verified using two separate checks. + +Digest verification uses the file hashes in `metadata.json` to verify the +integrity of kernel files. It is performed for both remote and local +kernels. If the files do not match the digest, an exception is raised. + +Digest verification only checks the integrity of the files according to +the metadata. Signature verification checks the authenticity of the metadata +itself: - Verify the signature against the given policy. The default policy only accepts kernels signed by workflows in the `huggingface/kernels-community` GitHub repository. - Verify the authenticity of `metadata.json` using the signature. -- Use the digests in `metadata.json` to verify the kernel files. +Signature verification is only performed for kernels downloaded from the Hub. At this time, a signature verification error will only result in a warning. Moreover, signature verification is only performed when the `sigstore` Python package is installed. However, we will make signature verification mandatory in the future. -The same steps can be performed on demand with the +Both checks can be performed on demand with the [`kernels verify-signature`](cli-verify-signature.md) command. -#### Signature verification receipts +#### Verification receipts -To avoid the high cost of signature verification, a kernel is only verified in -full once. The first time a kernel is loaded, we perform all the steps above. -Upon successful verification, we write a receipt file to the kernels cache. +To avoid the high cost of verification, a kernel from the Hub is only verified +in full once. Upon successful verification, we write a receipt file to the +kernels cache. Signature and digest verification each have their own receipts, +stored in `.verified-kernels/signature` and `.verified-kernels/digest`. -When a receipt is found on a later load, the signature and digest checks are -skipped. The signing certificate is still checked against the policy, since the +When a digest receipt is found on a later load, the kernel files are not hashed +again. When a signature receipt is found, the signature is not verified again. +However, the signing certificate is still checked against the policy, since the receipt could have been written by a verification with a different policy. +Local kernels do not get receipts, since their files may change. They are +hashed on every load. + Receipts are stored by kernel identity. The name of a receipt file is a hash of: - The repo ID diff --git a/kernels/src/kernels/cli/verify_signature.py b/kernels/src/kernels/cli/verify_signature.py index 5f442fd3..4293bd67 100644 --- a/kernels/src/kernels/cli/verify_signature.py +++ b/kernels/src/kernels/cli/verify_signature.py @@ -8,11 +8,13 @@ else: from typing_extensions import assert_never -from kernels._rust import KernelLocation +from kernels._rust import KernelLocation, Metadata from kernels._versions import resolve_revision_or_version +from kernels.digest import DigestVerificationResult, verify_digest from kernels.install import install_kernel, install_kernel_all_variants from kernels.variants import get_variants_local -from kernels.verify import VerificationResult, verify_variant +from kernels.verify import SignatureVerificationResult +from kernels.verify import verify_signature as verify_variant_signature def verify_signature(args: argparse.Namespace) -> None: @@ -34,23 +36,39 @@ def verify_signature(args: argparse.Namespace) -> None: for kernel_path in kernel_paths: variant_str = kernel_path.name + location = KernelLocation.remote(args.repo_id, revision, variant_str) - result = verify_variant( + signature_result = verify_variant_signature( kernel_path, - location=KernelLocation.remote(args.repo_id, revision, variant_str), + location=location, # Always fully verify the kernel in this subcommand. cache=False, ) - match result: - case VerificationResult.SignatureBundleMissing() if args.filter_unsigned: + match signature_result: + case SignatureVerificationResult.SignatureBundleMissing() if args.filter_unsigned: + continue + case SignatureVerificationResult.MetadataMissing() if args.filter_no_digest: + continue + case SignatureVerificationResult.Success(): pass - case VerificationResult.MetadataMissing() | VerificationResult.DigestMissing() if args.filter_no_digest: + case SignatureVerificationResult.Failure(): + print(f"❌ {variant_str}: {signature_result}") + failed = True + continue + case _ as unreachable: + assert_never(unreachable) + + metadata = Metadata.read_from_file(kernel_path / "metadata.json") + digest_result = verify_digest(kernel_path, metadata=metadata, location=location, cache=False) + + match digest_result: + case DigestVerificationResult.DigestMissing() if args.filter_no_digest: pass - case VerificationResult.Success(): - print(f"✅ {variant_str}: {result}") - case VerificationResult.Failure(): - print(f"❌ {variant_str}: {result}") + case DigestVerificationResult.Success(): + print(f"✅ {variant_str}: {signature_result} and {digest_result}") + case DigestVerificationResult.Failure(): + print(f"❌ {variant_str}: {digest_result}") failed = True case _ as unreachable: assert_never(unreachable) diff --git a/kernels/src/kernels/digest.py b/kernels/src/kernels/digest.py new file mode 100644 index 00000000..8ab61f0d --- /dev/null +++ b/kernels/src/kernels/digest.py @@ -0,0 +1,139 @@ +import abc +import logging +from dataclasses import dataclass +from pathlib import Path +from typing import TypeAlias, final + +from kernels._rust import ( + Digest, + DigestReceipt, + DigestReceiptStore, + DigestValidationError, + DigestViolation, + KernelLocation, + Metadata, + ReceiptError, + SignatureReceiptStore, +) + +logger = logging.getLogger(__name__) + + +class DigestVerificationResult: + class Failure(abc.ABC): + """A kernel build variant whose files could not be verified against its digest.""" + + @abc.abstractmethod + def __str__(self) -> str: ... + + @final + @dataclass + class DigestVerificationFailure(Failure): + """ + Verification failed because there were digest violations. + + The violations are provided through the `violations` field. + """ + + violations: list[DigestViolation] + + def __str__(self) -> str: + violations = "\n".join(str(violation) for violation in self.violations) + return f"the files do not match the digest in the metadata, so they may have been modified:\n{violations}" + + @final + @dataclass + class DigestMissing(Failure): + """ + Verification failed because the metadata did not have a digest. + """ + + def __str__(self) -> str: + return "the metadata does not record a digest, so its integrity cannot be verified" + + @final + @dataclass + class Success: + """ + Verification was successful. + """ + + def __str__(self) -> str: + return "the files match the digest in the metadata" + + Any: TypeAlias = DigestMissing | DigestVerificationFailure | Success + + +def _open_digest_receipt_store() -> DigestReceiptStore | None: + """The digest receipt store, or `None` when verifications cannot be cached.""" + try: + return DigestReceiptStore.in_kernels_cache() + except ReceiptError as e: + logger.warning(f"Cannot cache kernel digest verifications: {e}") + return None + + +def _has_receipt(store: DigestReceiptStore | SignatureReceiptStore, location: KernelLocation) -> bool: + """Whether the kernel at `location` was verified before. + + An unusable receipt counts as a cache miss: the kernel is then verified in + full, which overwrites the receipt. A broken cache must never make a kernel + fail to verify. + """ + try: + return store.load(location) is not None + except ReceiptError as e: + logger.warning(f"Ignoring unusable kernel verification receipt: {e}") + return False + + +def verify_digest( + variant_path: Path, + *, + metadata: Metadata, + location: KernelLocation | None, + cache: bool = True, +) -> DigestVerificationResult.Any: + """ + Verify the files of a kernel variant against the digest in its metadata. + + This only checks that the files and the metadata agree. It does not check + the authenticity of the metadata, use `kernels.verify.verify_signature` + for that. + + Args: + variant_path (`Path`): + Kernel variant path. + metadata (`Metadata`): + The metadata of the kernel variant. + location (`KernelLocation`, *optional*): + Identity of the kernel, used to cache the verification. Kernels + without a location (e.g. local kernels) are always verified in full. + cache (`bool`): + Whether to use the receipt cache to lookup or store kernel + verifications. + """ + if metadata.digest is None: + return DigestVerificationResult.DigestMissing() + + receipt_store = None + if cache and location is not None: + receipt_store = _open_digest_receipt_store() + if receipt_store is not None and _has_receipt(receipt_store, location): + return DigestVerificationResult.Success() + + # The validation is delegated to the (Rust) `Digest.validate`, which raises + # with each individual violation. + current_digest = Digest.hash_variant(metadata.digest.algorithm, variant_path) + try: + metadata.digest.validate(current_digest) + except DigestValidationError as e: + return DigestVerificationResult.DigestVerificationFailure(violations=e.violations) + + if receipt_store is not None and location is not None: + try: + receipt_store.store(DigestReceipt(location)) + except ReceiptError as e: + logger.warning(f"Cannot store kernel digest verification receipt: {e}") + + return DigestVerificationResult.Success() diff --git a/kernels/src/kernels/validate.py b/kernels/src/kernels/validate.py index 9fb0a3dd..4825ea28 100644 --- a/kernels/src/kernels/validate.py +++ b/kernels/src/kernels/validate.py @@ -10,6 +10,7 @@ else: from typing_extensions import assert_never +from kernels import digest from kernels._rust import KernelLocation, Metadata, Version from kernels.archs import _check_arch_incompatibility from kernels.backends import _backend @@ -157,6 +158,10 @@ class SignatureValidator: Only kernels with a known Hub origin are verified, since local kernels are typically for development and not signed. + This only verifies the authenticity of the kernel metadata. Use + [`DigestValidator`] to verify that the kernel file hashes match the + metadata. + Verification issues are currently reported as warnings. However, an exception will be raised in future versions.""" @@ -170,7 +175,7 @@ def validate_kernel(self, *, kernel: "LocalKernel") -> None: return # sigstore is still an optional dependency, so import lazily. - from kernels.verify import VerificationResult, verify_variant + from kernels.verify import SignatureVerificationResult, verify_signature location = KernelLocation.remote( kernel.origin.repo_id, @@ -178,22 +183,59 @@ def validate_kernel(self, *, kernel: "LocalKernel") -> None: kernel.variant_str, ) - result = verify_variant(kernel.variant_path, policy=self.policy, location=location) + result = verify_signature(kernel.variant_path, policy=self.policy, location=location) + + kernel_str = f"Kernel '{kernel.metadata.name}' variant '{kernel.variant_str}'" + + match result: + case SignatureVerificationResult.Success(): + logger.debug(f"{kernel_str}: {result}") + case SignatureVerificationResult.Failure(): + logger.warning(f"{kernel_str}: {result}", stacklevel=3) + case _ as unreachable: + assert_never(unreachable) + + +@dataclass +class DigestValidator: + """Verify that the files of a kernel build variant match the digest in its metadata. + + Kernels with a known Hub origin are only hashed once, later loads use the + verification receipt. Local kernels are hashed on every load, since their + files may change. + + Raises an exception when the files do not match the digest. A kernel without + a digest is loaded with a warning, since its integrity cannot be verified.""" + + def validate_kernel(self, *, kernel: "LocalKernel") -> None: + location = ( + KernelLocation.remote( + kernel.origin.repo_id, + kernel.origin.revision, + kernel.variant_str, + ) + if kernel.origin is not None + else None + ) + + result = digest.verify_digest(kernel.variant_path, metadata=kernel.metadata, location=location) kernel_str = f"Kernel '{kernel.metadata.name}' variant '{kernel.variant_str}'" match result: - case VerificationResult.Success(): + case digest.DigestVerificationResult.Success(): logger.debug(f"{kernel_str}: {result}") - case VerificationResult.Failure(): + case digest.DigestVerificationResult.DigestMissing(): logger.warning(f"{kernel_str}: {result}", stacklevel=3) + case digest.DigestVerificationResult.DigestVerificationFailure(): + raise RuntimeError(f"{kernel_str}: {result}") case _ as unreachable: assert_never(unreachable) def default_kernel_validators() -> list[KernelValidator]: """The kernel validators that are applied to every kernel dependency tree.""" - return [SignatureValidator()] + return [SignatureValidator(), DigestValidator()] @dataclass @@ -218,4 +260,8 @@ def validate_kernel(self, *, kernel: "LocalKernel") -> None: MinverValidator(), AllMetadataValidator([]), ) - _kernel_validator: tuple[KernelValidator, ...] = (SignatureValidator(), AllKernelValidator([])) + _kernel_validator: tuple[KernelValidator, ...] = ( + SignatureValidator(), + DigestValidator(), + AllKernelValidator([]), + ) diff --git a/kernels/src/kernels/verify.py b/kernels/src/kernels/verify.py index cd5af504..25b63ab2 100644 --- a/kernels/src/kernels/verify.py +++ b/kernels/src/kernels/verify.py @@ -11,15 +11,13 @@ from sigstore.verify.policy import VerificationPolicy from kernels._rust import ( - Digest, - DigestValidationError, - DigestViolation, KernelLocation, Metadata, ReceiptError, SignatureReceipt, SignatureReceiptStore, ) +from kernels.digest import _has_receipt logger = logging.getLogger(__name__) @@ -82,31 +80,13 @@ def verify(self, cert: Certificate) -> None: """ -class VerificationResult: +class SignatureVerificationResult: class Failure(abc.ABC): - """A kernel build variant that could not be verified.""" + """A kernel build variant whose signature could not be verified.""" @abc.abstractmethod def __str__(self) -> str: ... - @final - @dataclass - class DigestVerificationFailure(Failure): - """ - Verification failed because there were digest violations. - - The violations are provided through the `violations` field. - """ - - 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(Failure): @@ -143,16 +123,6 @@ class SignatureVerificationFailure(Failure): def __str__(self) -> str: return f"the metadata could not be verified against its signature:\n{self.reason}" - @final - @dataclass - class DigestMissing(Failure): - """ - Verification failed because the metadata did not have a digest. - """ - - def __str__(self) -> str: - return "the metadata does not record a digest, so its integrity cannot be verified" - @final @dataclass class MetadataMissing(Failure): @@ -184,9 +154,7 @@ def __str__(self) -> str: return "the metadata is correctly signed" Any: TypeAlias = ( - DigestMissing - | DigestVerificationFailure - | MetadataInvalid + MetadataInvalid | MetadataMissing | SignatureBundleInvalid | SignatureBundleMissing @@ -195,43 +163,31 @@ def __str__(self) -> str: ) -def _open_receipt_store() -> SignatureReceiptStore | None: - """The receipt store, or `None` when verifications cannot be cached.""" +def _open_signature_receipt_store() -> SignatureReceiptStore | None: + """The signature receipt store, or `None` when verifications cannot be cached.""" try: return SignatureReceiptStore.in_kernels_cache() except ReceiptError as e: - logger.warning(f"Cannot cache kernel verifications: {e}") + logger.warning(f"Cannot cache kernel signature verifications: {e}") return None -def _has_receipt(store: SignatureReceiptStore, location: KernelLocation) -> bool: - """Whether the kernel at `location` was verified before. - - An unusable receipt counts as a cache miss: the kernel is then verified in - full, which overwrites the receipt. A broken cache must never make a kernel - fail to verify. - """ - try: - return store.load(location) is not None - except ReceiptError as e: - logger.warning(f"Ignoring unusable kernel verification receipt: {e}") - return False - - -def verify_variant( +def verify_signature( variant_path: Path, *, location: KernelLocation, policy: VerificationPolicy | None = None, cache: bool = True, -) -> VerificationResult.Any: +) -> SignatureVerificationResult.Any: """ - Verify a kernel variant. + Verify the signature of a kernel variant. 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. + the given policy. + + This does not check that the files of the kernel match the digest in the + metadata, use `kernels.digest.verify_digest` for that. Args: variant_path (`Path`): @@ -252,32 +208,31 @@ def verify_variant( bundle_path = variant_path / "metadata.json.sigstore" if not bundle_path.is_file(): - return VerificationResult.SignatureBundleMissing() + return SignatureVerificationResult.SignatureBundleMissing() try: signature_bundle = Bundle.from_json(bundle_path.read_bytes()) except InvalidBundle as e: - return VerificationResult.SignatureBundleInvalid(reason=str(e)) + return SignatureVerificationResult.SignatureBundleInvalid(reason=str(e)) metadata_path = variant_path / "metadata.json" if not metadata_path.is_file(): - return VerificationResult.MetadataMissing() + return SignatureVerificationResult.MetadataMissing() - receipt_store = _open_receipt_store() if cache else None + receipt_store = _open_signature_receipt_store() if cache else None if receipt_store is not None and _has_receipt(receipt_store, location): - # The receipt attests that this kernel metadata was verified - # using the signature and the kernel data during the digest - # in the metadata. However, it may have been verified with a - # different policy, so we have to check certificate in the - # bundle against the currently required policy. + # The receipt attests that this kernel metadata was verified using + # the signature. However, it may have been verified with a different + # policy, so we have to check certificate in the bundle against the + # currently required policy. try: verify_policy.verify(signature_bundle.signing_certificate) except VerificationError as e: - return VerificationResult.SignatureVerificationFailure(reason=str(e)) + return SignatureVerificationResult.SignatureVerificationFailure(reason=str(e)) - return VerificationResult.Success() + return SignatureVerificationResult.Success() verifier = Verifier.production() @@ -292,30 +247,17 @@ def verify_variant( verify_policy, ) except VerificationError as e: - return VerificationResult.SignatureVerificationFailure(reason=str(e)) + return SignatureVerificationResult.SignatureVerificationFailure(reason=str(e)) try: - metadata = Metadata.from_bytes(metadata_bytes) + Metadata.from_bytes(metadata_bytes) except (OSError, ValueError) as e: - return VerificationResult.MetadataInvalid(reason=str(e)) - - if metadata.digest is None: - return VerificationResult.DigestMissing() - - hash = metadata.digest.algorithm - - # Rehash and check that the hashes match up. The validation is delegated to - # the (Rust) `Digest.validate`, which raises with each individual violation. - current_digest = Digest.hash_variant(hash, variant_path) - try: - metadata.digest.validate(current_digest) - except DigestValidationError as e: - return VerificationResult.DigestVerificationFailure(violations=e.violations) + return SignatureVerificationResult.MetadataInvalid(reason=str(e)) if receipt_store is not None: try: receipt_store.store(SignatureReceipt(location)) except ReceiptError as e: - logger.warning(f"Cannot store kernel verification receipt: {e}") + logger.warning(f"Cannot store kernel signature verification receipt: {e}") - return VerificationResult.Success() + return SignatureVerificationResult.Success() diff --git a/kernels/tests/test_digest.py b/kernels/tests/test_digest.py new file mode 100644 index 00000000..3be5e100 --- /dev/null +++ b/kernels/tests/test_digest.py @@ -0,0 +1,283 @@ +import json +import logging +import subprocess +import sys +from dataclasses import is_dataclass + +import pytest + +import kernels.digest as digest_module +from kernels._rust import ( + Digest, + DigestAlgorithm, + DigestReceiptStore, + DigestViolation, + KernelLocation, + Metadata, + Oid, +) +from kernels.digest import DigestVerificationResult, verify_digest + +_VARIANT = "torch-cpu" + + +def _write_variant(tmp_path, *, with_digest: bool = True): + """Write a kernel variant, with a digest of its files in the metadata.""" + variant_path = tmp_path / _VARIANT + variant_path.mkdir() + (variant_path / "__init__.py").write_text("from ._ops import ops\n") + (variant_path / "_ops.py").write_text("ops = None\n") + (variant_path / "_kernel.abi3.so").write_bytes(b"\x7fELF kernel") + + metadata = { + "name": "test-kernel", + "id": "_test_kernel_cpu_abc123", + "version": 1, + "license": "mit", + "python-depends": [], + "backend": {"type": "cpu"}, + } + if with_digest: + metadata["digest"] = { + "algorithm": "sha256", + "files": Digest.hash_variant(DigestAlgorithm.SHA256, variant_path).files, + } + (variant_path / "metadata.json").write_text(json.dumps(metadata)) + + return variant_path + + +def _metadata(variant_path) -> Metadata: + return Metadata.read_from_file(variant_path / "metadata.json") + + +def _location(variant_path) -> KernelLocation: + return KernelLocation.remote("kernels-test/digest", Oid.from_str("a" * 40), variant_path.name) + + +@pytest.fixture +def receipt_store(tmp_path, monkeypatch): + """An isolated receipt store, so that tests do not share verifications.""" + store = DigestReceiptStore.from_path(tmp_path / "receipts") + monkeypatch.setattr(digest_module, "_open_digest_receipt_store", lambda: store) + return store + + +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.digest`: `Digest` is an extension type, whose + # attributes cannot be set. + monkeypatch.setattr(digest_module, "Digest", ExplodingDigest) + + +def test_matching_files_pass(tmp_path): + variant_path = _write_variant(tmp_path) + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=None) + assert result == DigestVerificationResult.Success() + + +def test_modified_file_fails(tmp_path): + variant_path = _write_variant(tmp_path) + (variant_path / "_ops.py").write_text("ops = 'hacked'\n") + + match verify_digest(variant_path, metadata=_metadata(variant_path), location=None): + case DigestVerificationResult.DigestVerificationFailure( + violations=[DigestViolation.HashMismatch() as violation] + ): + assert violation.path == "_ops.py" + case other: + raise RuntimeError(f"Expected a single hash mismatch, was: {other}") + + +def test_added_file_fails(tmp_path): + variant_path = _write_variant(tmp_path) + (variant_path / "extra.py").write_text("print('hi')\n") + + match verify_digest(variant_path, metadata=_metadata(variant_path), location=None): + case DigestVerificationResult.DigestVerificationFailure( + violations=[DigestViolation.UnknownFile() as violation] + ): + assert violation.path == "extra.py" + case other: + raise RuntimeError(f"Expected a single unknown file, was: {other}") + + +def test_removed_file_fails(tmp_path): + variant_path = _write_variant(tmp_path) + (variant_path / "_ops.py").unlink() + + match verify_digest(variant_path, metadata=_metadata(variant_path), location=None): + case DigestVerificationResult.DigestVerificationFailure( + violations=[DigestViolation.MissingFile() as violation] + ): + assert violation.path == "_ops.py" + case other: + raise RuntimeError(f"Expected a single missing file, was: {other}") + + +def test_bytecode_is_ignored(tmp_path): + variant_path = _write_variant(tmp_path) + pycache = variant_path / "__pycache__" + pycache.mkdir() + (pycache / "_ops.cpython-314.pyc").write_bytes(b"bytecode") + + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=None) + assert result == DigestVerificationResult.Success() + + +def test_missing_digest(tmp_path): + variant_path = _write_variant(tmp_path, with_digest=False) + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=None) + assert result == DigestVerificationResult.DigestMissing() + + +def test_verification_is_cached(tmp_path, receipt_store, monkeypatch): + variant_path = _write_variant(tmp_path) + location = _location(variant_path) + + assert verify_digest(variant_path, metadata=_metadata(variant_path), location=location) == ( + DigestVerificationResult.Success() + ) + assert receipt_store.load(location) is not None + + # The second verification must be served from the receipt, without + # rehashing the variant. + _no_hashing(monkeypatch) + assert verify_digest(variant_path, metadata=_metadata(variant_path), location=location) == ( + DigestVerificationResult.Success() + ) + + +def test_failed_verification_is_not_cached(tmp_path, receipt_store): + variant_path = _write_variant(tmp_path) + (variant_path / "_ops.py").write_text("ops = 'hacked'\n") + location = _location(variant_path) + + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=location) + assert isinstance(result, DigestVerificationResult.DigestVerificationFailure) + assert receipt_store.load(location) is None + + +def test_verification_is_not_cached_with_cache_off(tmp_path, receipt_store, monkeypatch): + variant_path = _write_variant(tmp_path) + location = _location(variant_path) + + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=location, cache=False) + assert result == DigestVerificationResult.Success() + + # Nothing was recorded, ... + assert receipt_store.load(location) is None + + # ... and a verification with caching off does the full work even when a + # receipt does exist. + assert verify_digest(variant_path, metadata=_metadata(variant_path), location=location) == ( + DigestVerificationResult.Success() + ) + assert receipt_store.load(location) is not None + + _no_hashing(monkeypatch) + with pytest.raises(AssertionError, match="was rehashed"): + verify_digest(variant_path, metadata=_metadata(variant_path), location=location, cache=False) + + +def test_kernel_without_location_does_not_use_receipts(tmp_path, monkeypatch): + variant_path = _write_variant(tmp_path) + + def no_store(): + raise AssertionError("the receipt store was opened for a kernel without a location") + + monkeypatch.setattr(digest_module, "_open_digest_receipt_store", no_store) + + result = verify_digest(variant_path, metadata=_metadata(variant_path), location=None) + assert result == DigestVerificationResult.Success() + + +def test_unusable_receipt_falls_back_to_verification(tmp_path, receipt_store, caplog): + variant_path = _write_variant(tmp_path) + location = _location(variant_path) + + assert verify_digest(variant_path, metadata=_metadata(variant_path), location=location) == ( + DigestVerificationResult.Success() + ) + + (receipt_path,) = list((tmp_path / "receipts").iterdir()) + receipt_path.write_text("not a receipt") + + with caplog.at_level(logging.WARNING, logger="kernels.digest"): + assert verify_digest(variant_path, metadata=_metadata(variant_path), location=location) == ( + DigestVerificationResult.Success() + ) + + assert "unusable kernel verification receipt" in caplog.text + + +def test_digest_verification_does_not_need_sigstore(): + """Digest verification must work when the optional `sigstore` is not installed.""" + script = ( + "import sys\n" + # A `None` entry makes importing the module fail. + "sys.modules['sigstore'] = None\n" + "import kernels.digest, kernels.validate\n" + "from kernels.compat import has_sigstore\n" + "assert not has_sigstore\n" + ) + subprocess.run([sys.executable, "-c", script], check=True) + + +ALL_RESULTS = [ + DigestVerificationResult.Success(), + DigestVerificationResult.DigestMissing(), + DigestVerificationResult.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(DigestVerificationResult).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, DigestVerificationResult.Success) + assert isinstance(result, DigestVerificationResult.Failure) != is_success + + +def test_result_messages_include_their_detail(): + violations = [DigestViolation.MissingFile("kernel.py"), DigestViolation.UnknownFile("extra.so")] + message = str(DigestVerificationResult.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(DigestVerificationResult.Failure): + pass + + with pytest.raises(TypeError, match="abstract"): + Undescribed() # type: ignore[abstract] diff --git a/kernels/tests/test_validate.py b/kernels/tests/test_validate.py index 97726fa1..280345d0 100644 --- a/kernels/tests/test_validate.py +++ b/kernels/tests/test_validate.py @@ -7,21 +7,25 @@ import torch import kernels +import kernels.digest as digest_module import kernels.validate as validate_module import kernels.verify as verify_module -from kernels._rust import KernelLocation, Metadata, Oid, Version +from kernels._rust import DigestViolation, KernelLocation, Metadata, Oid, Version from kernels.deps import DepTreeNode +from kernels.digest import DigestVerificationResult from kernels.resolver import LocalKernel, RemoteKernel from kernels.validate import ( ArchValidator, + DigestValidator, DirtyValidator, MinverValidator, SignatureValidator, _installed_version, + default_kernel_validators, default_metadata_validators, ) from kernels.variants import parse_variant -from kernels.verify import VerificationResult +from kernels.verify import SignatureVerificationResult CLEAN_PROVENANCE = { "kernel-builder": {"version": "0.1.0", "commit": "a" * 40, "dirty": False}, @@ -233,15 +237,15 @@ def _hub_kernel(tmp_path, metadata) -> LocalKernel: @pytest.fixture def recorded_verifications(monkeypatch): - """Record `verify_variant` calls and control the result it returns.""" + """Record `verify_signature` calls and control the result it returns.""" calls = [] results = [] - def fake_verify_variant(variant_path, *, location, policy=None, cache=True): + def fake_verify_signature(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() + return results.pop(0) if results else SignatureVerificationResult.Success() - monkeypatch.setattr(verify_module, "verify_variant", fake_verify_variant) + monkeypatch.setattr(verify_module, "verify_signature", fake_verify_signature) return calls, results @@ -290,13 +294,11 @@ def test_signature_validator_is_quiet_on_success(tmp_path, make_metadata, record @pytest.mark.parametrize( "result", [ - VerificationResult.SignatureBundleMissing(), - VerificationResult.SignatureBundleInvalid(reason="bad bundle"), - VerificationResult.SignatureVerificationFailure(reason="bad signature"), - VerificationResult.MetadataInvalid(reason="bad metadata"), - VerificationResult.MetadataMissing(), - VerificationResult.DigestMissing(), - VerificationResult.DigestVerificationFailure(violations=[]), + SignatureVerificationResult.SignatureBundleMissing(), + SignatureVerificationResult.SignatureBundleInvalid(reason="bad bundle"), + SignatureVerificationResult.SignatureVerificationFailure(reason="bad signature"), + SignatureVerificationResult.MetadataInvalid(reason="bad metadata"), + SignatureVerificationResult.MetadataMissing(), ], ) def test_signature_validator_warns_but_does_not_raise(tmp_path, make_metadata, recorded_verifications, caplog, result): @@ -311,3 +313,95 @@ def test_signature_validator_warns_but_does_not_raise(tmp_path, make_metadata, r # it applies to, so the wording is asserted where it is defined. assert str(result) in caplog.text assert "test-kernel" in caplog.text + + +@pytest.fixture +def recorded_digest_verifications(monkeypatch): + """Record `verify_digest` calls and control the result it returns.""" + calls = [] + results = [] + + def fake_verify_digest(variant_path, *, metadata, location, cache=True): + calls.append({"variant_path": variant_path, "metadata": metadata, "location": location, "cache": cache}) + return results.pop(0) if results else DigestVerificationResult.Success() + + monkeypatch.setattr(digest_module, "verify_digest", fake_verify_digest) + return calls, results + + +def test_digest_validator_verifies_local_kernels_without_location( + tmp_path, make_metadata, recorded_digest_verifications +): + calls, _ = recorded_digest_verifications + metadata = make_metadata("cuda", None) + kernel = LocalKernel(variant_path=tmp_path / "torch-cuda", metadata=metadata) + + DigestValidator().validate_kernel(kernel=kernel) + + (call,) = calls + assert call["variant_path"] == kernel.variant_path + assert call["metadata"] is metadata + # Local kernels may change, so they are always hashed. + assert call["location"] is None + + +def test_digest_validator_identifies_hub_kernel_by_origin(tmp_path, make_metadata, recorded_digest_verifications): + calls, _ = recorded_digest_verifications + kernel = _hub_kernel(tmp_path, make_metadata("cuda", None)) + + DigestValidator().validate_kernel(kernel=kernel) + + (call,) = calls + assert call["location"] == KernelLocation.remote(_SIGNED_REPO_ID, _SIGNED_REVISION, "torch-cuda") + # Loading a kernel must reuse a previous verification. + assert call["cache"] is True + + +def test_digest_validator_does_not_need_sigstore(tmp_path, make_metadata, recorded_digest_verifications, monkeypatch): + calls, _ = recorded_digest_verifications + monkeypatch.setattr(validate_module, "has_sigstore", False) + kernel = _hub_kernel(tmp_path, make_metadata("cuda", None)) + + DigestValidator().validate_kernel(kernel=kernel) + + assert len(calls) == 1 + + +def test_digest_validator_is_quiet_on_success(tmp_path, make_metadata, recorded_digest_verifications, caplog): + kernel = _hub_kernel(tmp_path, make_metadata("cuda", None)) + + with caplog.at_level(logging.WARNING, logger="kernels.validate"): + DigestValidator().validate_kernel(kernel=kernel) + + assert caplog.text == "" + + +def test_digest_validator_warns_on_missing_digest(tmp_path, make_metadata, recorded_digest_verifications, caplog): + _, results = recorded_digest_verifications + result = DigestVerificationResult.DigestMissing() + results.append(result) + kernel = _hub_kernel(tmp_path, make_metadata("cuda", None)) + + with caplog.at_level(logging.WARNING, logger="kernels.validate"): + DigestValidator().validate_kernel(kernel=kernel) + + assert str(result) in caplog.text + assert "test-kernel" in caplog.text + + +def test_digest_validator_raises_on_mismatch(tmp_path, make_metadata, recorded_digest_verifications): + _, results = recorded_digest_verifications + result = DigestVerificationResult.DigestVerificationFailure(violations=[DigestViolation.MissingFile("kernel.py")]) + results.append(result) + kernel = _hub_kernel(tmp_path, make_metadata("cuda", None)) + + with pytest.raises(RuntimeError) as exc_info: + DigestValidator().validate_kernel(kernel=kernel) + + assert str(result) in str(exc_info.value) + assert "test-kernel" in str(exc_info.value) + + +def test_default_kernel_validators_check_signature_and_digest(): + validators = default_kernel_validators() + assert [type(validator) for validator in validators] == [SignatureValidator, DigestValidator] diff --git a/kernels/tests/test_verify.py b/kernels/tests/test_verify.py index 39f4a5ca..780b3973 100644 --- a/kernels/tests/test_verify.py +++ b/kernels/tests/test_verify.py @@ -7,11 +7,12 @@ import kernels.verify as verify_module from kernels import install_kernel -from kernels._rust import DigestViolation, KernelLocation, Oid, SignatureReceiptStore +from kernels._rust import DigestViolation, KernelLocation, Metadata, Oid, SignatureReceiptStore from kernels._versions import resolve_revision_or_version +from kernels.digest import DigestVerificationResult, verify_digest from kernels.hf_hub import _get_cache_dir, _get_hf_api from kernels.resolver import _BYTECODE_IGNORE_PATTERNS -from kernels.verify import VerificationResult, verify_variant +from kernels.verify import SignatureVerificationResult, verify_signature TEST_POLICY: policy.VerificationPolicy = policy.Identity( identity="me@danieldk.eu", issuer="https://github.com/login/oauth" @@ -27,7 +28,7 @@ def receipt_store(tmp_path, monkeypatch): """An isolated receipt store, so that tests do not share verifications.""" receipt_dir = tmp_path / "receipts" store = SignatureReceiptStore.from_path(receipt_dir) - monkeypatch.setattr(verify_module, "_open_receipt_store", lambda: store) + monkeypatch.setattr(verify_module, "_open_signature_receipt_store", lambda: store) return store @@ -40,14 +41,14 @@ def signed_kernel(): 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. +def _verify_signature_uncached(variant_path: Path, **kwargs) -> SignatureVerificationResult.Any: + """Verify the signature of 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( + return verify_signature( variant_path, location=KernelLocation.remote("kernels-test/signatures", Oid.from_str("0" * 40), variant_path.name), cache=False, @@ -55,36 +56,46 @@ def _verify_uncached(variant_path: Path, **kwargs) -> VerificationResult.Any: ) -def _no_hashing(monkeypatch): - """Make rehashing the variant fail, so that only cache hits can succeed.""" +def _verify_digest_uncached(variant_path: Path) -> DigestVerificationResult.Any: + """Verify the files of a variant against its digest, without receipts.""" + metadata = Metadata.read_from_file(variant_path / "metadata.json") + return verify_digest(variant_path, metadata=metadata, location=None, cache=False) - class ExplodingDigest: + +def _no_signature_verification(monkeypatch): + """Make full signature verification fail, so that only cache hits can succeed.""" + + class ExplodingVerifier: @staticmethod - def hash_variant(*args, **kwargs): - raise AssertionError("the variant was rehashed, so this was not a cache hit") + def production(*args, **kwargs): + raise AssertionError("the signature was verified in full, 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) + monkeypatch.setattr(verify_module, "Verifier", ExplodingVerifier) def test_correctly_signed_kernel_passes_with_default_policy(): 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() + assert _verify_signature_uncached(variant_path) == SignatureVerificationResult.Success() + assert _verify_digest_uncached(variant_path) == DigestVerificationResult.Success() def test_correctly_signed_kernel_passes(): 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() + assert _verify_signature_uncached(variant_path, policy=TEST_POLICY) == SignatureVerificationResult.Success() + assert _verify_digest_uncached(variant_path) == DigestVerificationResult.Success() def test_invalid_digest_fails(): variant_path = install_kernel("kernels-test/signatures", revision="invalid-digest") - match _verify_uncached(variant_path, policy=TEST_POLICY): - case VerificationResult.DigestVerificationFailure(violations=violations): + # The metadata itself is correctly signed, signature verification does + # not check the files. + assert _verify_signature_uncached(variant_path, policy=TEST_POLICY) == SignatureVerificationResult.Success() + + match _verify_digest_uncached(variant_path): + case DigestVerificationResult.DigestVerificationFailure(violations=violations): assert len(violations) == 1 assert isinstance(violations[0], DigestViolation.HashMismatch) case other: @@ -117,12 +128,12 @@ def test_invalid_metadata_fails(): / "build" ) - match _verify_uncached( + match _verify_signature_uncached( # No CUDA dependency, we are only checking metadata. variant_paths / "torch-cuda", policy=TEST_POLICY, ): - case VerificationResult.MetadataInvalid(reason=reason): + case SignatureVerificationResult.MetadataInvalid(reason=reason): assert "Cannot parse metadata" in reason case other: raise RuntimeError(f"Expected MetadataInvalid, was: {other}") @@ -130,7 +141,8 @@ def test_invalid_metadata_fails(): def test_missing_digest_fails(): variant_path = install_kernel("kernels-test/signatures", revision="missing-digest") - assert _verify_uncached(variant_path, policy=TEST_POLICY) == VerificationResult.DigestMissing() + assert _verify_signature_uncached(variant_path, policy=TEST_POLICY) == SignatureVerificationResult.Success() + assert _verify_digest_uncached(variant_path) == DigestVerificationResult.DigestMissing() def test_missing_metadata_fails(): @@ -160,24 +172,27 @@ def test_missing_metadata_fails(): ) assert ( - _verify_uncached( + _verify_signature_uncached( # No CUDA dependency, we are only checking metadata. variant_paths / "torch-cuda", policy=TEST_POLICY, ) - == VerificationResult.MetadataMissing() + == SignatureVerificationResult.MetadataMissing() ) def test_unsigned_kernel_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-missing") - assert _verify_uncached(variant_path, policy=TEST_POLICY) == VerificationResult.SignatureBundleMissing() + assert ( + _verify_signature_uncached(variant_path, policy=TEST_POLICY) + == SignatureVerificationResult.SignatureBundleMissing() + ) def test_broken_signature_bundle_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-broken") - match _verify_uncached(variant_path, policy=TEST_POLICY): - case VerificationResult.SignatureBundleInvalid(reason=_): + match _verify_signature_uncached(variant_path, policy=TEST_POLICY): + case SignatureVerificationResult.SignatureBundleInvalid(reason=_): pass case other: raise RuntimeError(f"Expected SignatureBundleInvalid, was: {other}") @@ -185,8 +200,8 @@ def test_broken_signature_bundle_fails(): def test_invalid_signature_fails(): variant_path = install_kernel("kernels-test/signatures", revision="signature-invalid") - match _verify_uncached(variant_path, policy=TEST_POLICY): - case VerificationResult.SignatureVerificationFailure(reason=_): + match _verify_signature_uncached(variant_path, policy=TEST_POLICY): + case SignatureVerificationResult.SignatureVerificationFailure(reason=_): pass case other: raise RuntimeError(f"Expected SignatureVerificationFailure, was: {other}") @@ -195,45 +210,53 @@ def test_invalid_signature_fails(): def test_verification_is_cached(receipt_store, signed_kernel, monkeypatch): variant_path, location = signed_kernel - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) == SignatureVerificationResult.Success() + ) assert receipt_store.load(location) is not None # The second verification must be served from the receipt, without - # rehashing the variant. - _no_hashing(monkeypatch) - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + # verifying the signature in full. + _no_signature_verification(monkeypatch) + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) == SignatureVerificationResult.Success() + ) def test_verification_is_not_cached_with_cache_off(receipt_store, signed_kernel, monkeypatch): variant_path, location = signed_kernel - result = verify_variant(variant_path, policy=TEST_POLICY, location=location, cache=False) - assert result == VerificationResult.Success() + result = verify_signature(variant_path, policy=TEST_POLICY, location=location, cache=False) + assert result == SignatureVerificationResult.Success() # Nothing was recorded, ... assert receipt_store.load(location) is None # ... and a verification with caching off does the full work even when a # receipt does exist. - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) == SignatureVerificationResult.Success() + ) assert receipt_store.load(location) is not None - _no_hashing(monkeypatch) - with pytest.raises(AssertionError, match="was rehashed"): - verify_variant(variant_path, policy=TEST_POLICY, location=location, cache=False) + _no_signature_verification(monkeypatch) + with pytest.raises(AssertionError, match="verified in full"): + verify_signature(variant_path, policy=TEST_POLICY, location=location, cache=False) def test_cached_verification_still_enforces_policy(receipt_store, signed_kernel, monkeypatch): variant_path, location = signed_kernel # Verify under a policy that accepts this kernel, so a receipt is stored. - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) == SignatureVerificationResult.Success() + ) # The receipt says the kernel was verified, but not *under which policy*, # so a policy that does not accept this signer must still reject it. - _no_hashing(monkeypatch) - match verify_variant(variant_path, policy=OTHER_POLICY, location=location): - case VerificationResult.SignatureVerificationFailure(): + _no_signature_verification(monkeypatch) + match verify_signature(variant_path, policy=OTHER_POLICY, location=location): + case SignatureVerificationResult.SignatureVerificationFailure(): pass case other: raise RuntimeError(f"Expected SignatureVerificationFailure, was: {other}") @@ -242,26 +265,29 @@ def test_cached_verification_still_enforces_policy(receipt_store, signed_kernel, def test_unusable_receipt_falls_back_to_verification(receipt_store, signed_kernel, tmp_path, caplog): variant_path, location = signed_kernel - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) == SignatureVerificationResult.Success() + ) (receipt_path,) = list((tmp_path / "receipts").iterdir()) receipt_path.write_text("not a receipt") - with caplog.at_level(logging.WARNING, logger="kernels.verify"): - assert verify_variant(variant_path, policy=TEST_POLICY, location=location) == VerificationResult.Success() + with caplog.at_level(logging.WARNING): + assert ( + verify_signature(variant_path, policy=TEST_POLICY, location=location) + == SignatureVerificationResult.Success() + ) assert "unusable kernel verification receipt" in caplog.text 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")]), + SignatureVerificationResult.Success(), + SignatureVerificationResult.SignatureBundleMissing(), + SignatureVerificationResult.SignatureBundleInvalid(reason="bad bundle"), + SignatureVerificationResult.SignatureVerificationFailure(reason="bad signature"), + SignatureVerificationResult.MetadataMissing(), + SignatureVerificationResult.MetadataInvalid(reason="bad metadata"), ] @@ -272,7 +298,9 @@ def test_all_results_are_covered(): is exactly when they are needed. """ variants = { - name for name, member in vars(VerificationResult).items() if isinstance(member, type) and is_dataclass(member) + name + for name, member in vars(SignatureVerificationResult).items() + if isinstance(member, type) and is_dataclass(member) } assert {type(result).__name__ for result in ALL_RESULTS} == variants @@ -288,25 +316,20 @@ def test_every_result_describes_itself(result): @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 + is_success = isinstance(result, SignatureVerificationResult.Success) + assert isinstance(result, SignatureVerificationResult.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 + assert "bang" in str(SignatureVerificationResult.SignatureBundleInvalid(reason="bang")) + assert "bang" in str(SignatureVerificationResult.MetadataInvalid(reason="bang")) + assert "bang" in str(SignatureVerificationResult.SignatureVerificationFailure(reason="bang")) def test_failure_must_describe_itself(): """The base class makes a message mandatory for new failures.""" - class Undescribed(VerificationResult.Failure): + class Undescribed(SignatureVerificationResult.Failure): pass with pytest.raises(TypeError, match="abstract"): From 7e53283fff7bd4d89c9e6241eee079eabcd801e9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Fri, 9 Oct 2026 15:27:20 +0000 Subject: [PATCH 5/5] kernels: raise when two different builds have the same ID --- kernels/src/kernels/importer.py | 85 ++++++++++++--- kernels/tests/test_importer.py | 177 +++++++++++++++++++++++++++++++- 2 files changed, 244 insertions(+), 18 deletions(-) diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index ba32d72d..09d56ec9 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -1,11 +1,12 @@ import importlib import logging import sys +import threading from dataclasses import dataclass from pathlib import Path from types import ModuleType -from kernels._rust import Metadata +from kernels._rust import DigestValidationError, Metadata from kernels.hf_hub import RepoInfo logger = logging.getLogger(__name__) @@ -42,6 +43,14 @@ class LoadedKernel: _loaded_kernels: dict[Path, LoadedKernel] = {} +# Metadata of loaded kernels by their ids. +_loaded_kernel_metadata: dict[str, Metadata] = {} + +# Serializes kernel imports, so that a kernel is executed only once and other +# threads never see a partially initialized kernel module. Reentrant, so that a +# kernel that loads another kernel while it is imported does not deadlock. +_import_lock = threading.RLock() + def get_loaded_kernels() -> list[LoadedKernel]: """ @@ -65,20 +74,63 @@ def get_loaded_kernels() -> list[LoadedKernel]: return list(_loaded_kernels.values()) +def _check_same_build(variant_path: Path, metadata: Metadata) -> None: + """Check that the kernel at `variant_path` is the same build as the loaded + kernel with the same id. + + If there are two different kernels with the same kernel id, one (or both) + of the kernels violates the unique kernel id requirement. + + Raises `RuntimeError` when the builds differ. + """ + reference_metadata = _loaded_kernel_metadata.get(metadata.id) + assert reference_metadata is not None, ( + f"Kernel '{metadata.id}' is in `sys.modules`, but its metadata was not recorded" + ) + + reference_digest = reference_metadata.digest + digest = metadata.digest + + if reference_digest is None and digest is None: + logger.debug( + f"Cannot compare kernel '{metadata.id}' at `{variant_path}` with the loaded build: neither has a digest" + ) + return + + conflict = ( + f"Kernel '{metadata.name}' at `{variant_path}` has the same id '{metadata.id}' as a kernel " + "that is already loaded, but it is a different build" + ) + hint = "Was the kernel modified without rebuilding it with `kernel-builder`?" + + if reference_digest is None: + raise RuntimeError(f"{conflict}: this build has a digest, but the loaded build does not. {hint}") + if digest is None: + raise RuntimeError(f"{conflict}: the loaded build has a digest, but this build does not. {hint}") + + try: + reference_digest.validate(digest) + except DigestValidationError as e: + violations = "\n".join(str(violation) for violation in e.violations) + raise RuntimeError(f"{conflict}. {hint}\nDifferences with the loaded build:\n{violations}") from e + + def _import_from_path( variant_path: Path, deps: dict[str, ModuleType], repo_info: RepoInfo | None = None, ) -> ModuleType: - if (loaded_kernel := _loaded_kernels.get(variant_path)) is not None: - return loaded_kernel.module - metadata = Metadata.read_from_file(variant_path / "metadata.json") - # Kernel ids are unique per build: if this build was already imported - # reuse it instead of executing it again. - if (module := sys.modules.get(metadata.id)) is not None: - logging.debug(f"Kernel already loaded, skipping: {metadata.id}") + with _import_lock: + # Kernel ids are unique per build: if this build was already imported + # reuse it instead of executing it again. + if (module := sys.modules.get(metadata.id)) is None: + module = _import_from_path_uncached(variant_path, metadata, deps, repo_info) + else: + logger.debug(f"Kernel already loaded, skipping: {metadata.id}") + _check_same_build(variant_path, metadata) + _loaded_kernels[variant_path] = LoadedKernel( metadata=metadata, module=module, @@ -86,6 +138,17 @@ def _import_from_path( ) return module + +def _import_from_path_uncached( + variant_path: Path, + metadata: Metadata, + deps: dict[str, ModuleType], + repo_info: RepoInfo | None, +) -> ModuleType: + """Import the kernel at `variant_path`, without reusing an imported kernel with the same id. + + Must be called with `_import_lock` held. + """ module_name = metadata.name.python_name file_path = variant_path / "__init__.py" @@ -117,9 +180,5 @@ def _import_from_path( e.add_note(f"while importing kernel '{metadata.name}', variant '{variant_path.name}' {origin}") raise - _loaded_kernels[variant_path] = LoadedKernel( - metadata=metadata, - module=module, - repo_info=repo_info, - ) + _loaded_kernel_metadata[metadata.id] = metadata return module diff --git a/kernels/tests/test_importer.py b/kernels/tests/test_importer.py index 7fd53545..87d8ce22 100644 --- a/kernels/tests/test_importer.py +++ b/kernels/tests/test_importer.py @@ -1,13 +1,16 @@ import json import sys +import threading import types import pytest -from kernels.importer import _import_from_path, _loaded_kernels +from kernels._rust import Digest, DigestAlgorithm +from kernels.importer import _import_from_path, _loaded_kernel_metadata, _loaded_kernels _EXEC_LOG_MODULE = "_kernels_test_exec_log" _COUNTING_ID = "counting_1_cuda" +_REBUILT_ID = "counting_2_cuda" def _write_variant(tmp_path): @@ -38,13 +41,33 @@ def test_failed_import_cleans_up_sys_modules(tmp_path): assert any("broken" in note for note in exc_info.value.__notes__) finally: _loaded_kernels.pop(variant_dir, None) + _loaded_kernel_metadata.pop("broken_1_cuda", None) sys.modules.pop("broken_1_cuda", None) -def _write_counting_variant(base_path, kernel_id): - """Write a kernel variant that records each execution of its module.""" +def _write_counting_variant( + base_path, + kernel_id, + *, + variant_tag: str = "", + with_digest: bool = False, + init_delay: float = 0.0, +): + """Write a kernel variant that records each execution of its module. + + Variants with a different `variant_tag` have different files. The module + sleeps for `init_delay` seconds before it is fully initialized, which is + the case once `TAG` is set. + """ variant_dir = base_path / "build" / "torch28-cxx11-cu128-x86_64-linux" - variant_dir.mkdir(parents=True) + variant_dir.mkdir(parents=True, exist_ok=True) + (variant_dir / "__init__.py").write_text( + "import time\n" + f"import {_EXEC_LOG_MODULE}\n" + f"{_EXEC_LOG_MODULE}.calls.append(__file__)\n" + f"time.sleep({init_delay!r})\n" + f"TAG = {variant_tag!r}\n" + ) metadata = { "id": kernel_id, "name": "counting", @@ -53,8 +76,12 @@ def _write_counting_variant(base_path, kernel_id): "python-depends": ["torch"], "backend": {"type": "cuda"}, } + if with_digest: + metadata["digest"] = { + "algorithm": "sha256", + "files": Digest.hash_variant(DigestAlgorithm.SHA256, variant_dir).files, + } (variant_dir / "metadata.json").write_text(json.dumps(metadata)) - (variant_dir / "__init__.py").write_text(f"import {_EXEC_LOG_MODULE}\n{_EXEC_LOG_MODULE}.calls.append(__file__)\n") return variant_dir @@ -66,6 +93,18 @@ def exec_log(monkeypatch): return log.calls # type: ignore[attr-defined] +@pytest.fixture +def counting_cleanup(): + """Remove every registration of the counting kernels after the test.""" + yield + kernel_ids = {_COUNTING_ID, _REBUILT_ID} + for path in [path for path, loaded in _loaded_kernels.items() if loaded.metadata.id in kernel_ids]: + _loaded_kernels.pop(path) + for kernel_id in kernel_ids: + _loaded_kernel_metadata.pop(kernel_id, None) + sys.modules.pop(kernel_id, None) + + def test_same_id_different_path_is_not_reloaded(tmp_path, exec_log): first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID) second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID) @@ -80,9 +119,136 @@ def test_same_id_different_path_is_not_reloaded(tmp_path, exec_log): finally: _loaded_kernels.pop(first_dir, None) _loaded_kernels.pop(second_dir, None) + _loaded_kernel_metadata.pop(_COUNTING_ID, None) sys.modules.pop(_COUNTING_ID, None) +def test_same_id_same_digest_is_reused(tmp_path, exec_log, counting_cleanup): + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID, with_digest=True) + second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID, with_digest=True) + + first = _import_from_path(first_dir, deps={}) + second = _import_from_path(second_dir, deps={}) + + assert first is second + assert len(exec_log) == 1 + assert _loaded_kernels[first_dir].module is first + assert _loaded_kernels[second_dir].module is first + + +def test_same_id_different_digest_raises(tmp_path, exec_log, counting_cleanup): + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID, with_digest=True) + second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID, variant_tag="hacked", with_digest=True) + + _import_from_path(first_dir, deps={}) + with pytest.raises(RuntimeError, match="different build") as exc_info: + _import_from_path(second_dir, deps={}) + + message = str(exc_info.value) + assert str(second_dir) in message + assert _COUNTING_ID in message + assert "__init__.py" in message + assert len(exec_log) == 1 + assert second_dir not in _loaded_kernels + + +def test_same_id_check_does_not_depend_on_loaded_kernels(tmp_path, exec_log, counting_cleanup): + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID, with_digest=True) + second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID, variant_tag="hacked", with_digest=True) + + _import_from_path(first_dir, deps={}) + # The registry of loaded kernel paths does not provide the reference build. + _loaded_kernels.pop(first_dir) + + with pytest.raises(RuntimeError, match="different build"): + _import_from_path(second_dir, deps={}) + + +@pytest.mark.parametrize( + ("first_has_digest", "second_has_digest"), + [(True, False), (False, True)], + ids=["only-first-has-digest", "only-second-has-digest"], +) +def test_same_id_one_sided_digest_raises(tmp_path, exec_log, counting_cleanup, first_has_digest, second_has_digest): + # The files are identical, only the presence of a digest differs. + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID, with_digest=first_has_digest) + second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID, with_digest=second_has_digest) + + _import_from_path(first_dir, deps={}) + with pytest.raises(RuntimeError, match="has a digest, but"): + _import_from_path(second_dir, deps={}) + + assert len(exec_log) == 1 + assert second_dir not in _loaded_kernels + + +def test_in_place_rebuild_with_same_id_raises(tmp_path, exec_log, counting_cleanup): + variant_dir = _write_counting_variant(tmp_path, _COUNTING_ID, with_digest=True) + _import_from_path(variant_dir, deps={}) + + # Modify the kernel in place, including its digest, but keep its id. + _write_counting_variant(tmp_path, _COUNTING_ID, variant_tag="hacked", with_digest=True) + + with pytest.raises(RuntimeError, match="different build"): + _import_from_path(variant_dir, deps={}) + assert len(exec_log) == 1 + + +def test_in_place_rebuild_with_new_id_is_loaded(tmp_path, exec_log, counting_cleanup): + variant_dir = _write_counting_variant(tmp_path, _COUNTING_ID, with_digest=True) + first = _import_from_path(variant_dir, deps={}) + + # A proper rebuild gets a new kernel id. + _write_counting_variant(tmp_path, _REBUILT_ID, variant_tag="rebuilt", with_digest=True) + second = _import_from_path(variant_dir, deps={}) + + assert second is not first + assert second.TAG == "rebuilt" + assert len(exec_log) == 2 + assert _loaded_kernels[variant_dir].module is second + + +def _load_concurrently(paths): + """Import the given kernel paths from one thread each, started at the same time.""" + barrier = threading.Barrier(len(paths)) + modules = [None] * len(paths) + errors = [] + + def load(index, path): + barrier.wait() + try: + modules[index] = _import_from_path(path, deps={}) + except BaseException as e: + errors.append(e) + + threads = [threading.Thread(target=load, args=(index, path)) for index, path in enumerate(paths)] + for thread in threads: + thread.start() + for thread in threads: + thread.join(timeout=30) + + assert not any(thread.is_alive() for thread in threads), "kernel import did not finish" + assert errors == [] + return modules + + +@pytest.mark.parametrize("same_path", [True, False], ids=["same-path", "same-id-other-path"]) +def test_concurrent_imports_execute_kernel_once(tmp_path, exec_log, counting_cleanup, same_path): + first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID, with_digest=True, init_delay=0.2) + second_dir = ( + first_dir + if same_path + else _write_counting_variant(tmp_path / "second", _COUNTING_ID, with_digest=True, init_delay=0.2) + ) + + first, second = _load_concurrently([first_dir, second_dir]) + + assert len(exec_log) == 1 + assert first is second + # Neither thread got a partially initialized module. + assert first.TAG == "" + + def test_already_imported_kernel_is_reregistered(tmp_path, exec_log): variant_dir = _write_counting_variant(tmp_path, _COUNTING_ID) try: @@ -96,4 +262,5 @@ def test_already_imported_kernel_is_reregistered(tmp_path, exec_log): assert _loaded_kernels[variant_dir].module is first finally: _loaded_kernels.pop(variant_dir, None) + _loaded_kernel_metadata.pop(_COUNTING_ID, None) sys.modules.pop(_COUNTING_ID, None)