diff --git a/kernels/src/kernels/importer.py b/kernels/src/kernels/importer.py index c18d930a..5d0db257 100644 --- a/kernels/src/kernels/importer.py +++ b/kernels/src/kernels/importer.py @@ -1,4 +1,6 @@ import importlib +import importlib.machinery +import os import sys from dataclasses import dataclass from pathlib import Path @@ -62,6 +64,39 @@ def get_loaded_kernels() -> list[LoadedKernel]: return list(_loaded_kernels.values()) +class _SourceOnlyLoader(importlib.machinery.SourceFileLoader): + """ + Source loader that always compiles from the `.py` file. + + Bytecode is excluded from the kernel digest, since the interpreter + writes it after the first import. So bytecode cannot be trusted and + should never be executed in place of the verified source. + """ + + def get_code(self, fullname): + source_path = self.get_filename(fullname) + return self.source_to_code(self.get_data(source_path), source_path) + + def set_data(self, path, data, *, _mode=0o666): + # Do not write bytecode that is never read. + pass + + +def _register_source_only_finders(module_dir: Path): + """ + Register finders for every directory in the kernel module, so that + submodules are also loaded from source. Sourceless `.pyc` files are + not importable at all. + """ + loaders = [ + (importlib.machinery.ExtensionFileLoader, importlib.machinery.EXTENSION_SUFFIXES), + (_SourceOnlyLoader, importlib.machinery.SOURCE_SUFFIXES), + ] + for root, dirs, _ in os.walk(module_dir): + dirs[:] = [d for d in dirs if d != "__pycache__"] + sys.path_importer_cache[root] = importlib.machinery.FileFinder(root, *loaders) + + def _import_from_path( variant_path: Path, deps: dict[str, ModuleType], @@ -79,7 +114,10 @@ def _import_from_path( if not file_path.exists(): raise FileNotFoundError(f"No kernel module found at: `{variant_path}`") - spec = importlib.util.spec_from_file_location(metadata.id, file_path) + _register_source_only_finders(file_path.parent) + spec = importlib.util.spec_from_file_location( + metadata.id, file_path, loader=_SourceOnlyLoader(metadata.id, str(file_path)) + ) if spec is None: raise ImportError(f"Cannot load spec for {module_name} from {file_path}") module = importlib.util.module_from_spec(spec) diff --git a/kernels/tests/test_importer.py b/kernels/tests/test_importer.py index 3acb0bbb..a3be080f 100644 --- a/kernels/tests/test_importer.py +++ b/kernels/tests/test_importer.py @@ -1,4 +1,6 @@ +import importlib import json +import py_compile import sys import pytest @@ -35,3 +37,38 @@ 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 _plant_pyc(source_path, payload, tmp_path): + # An unchecked-hash .pyc (PEP 552) is used by the default loader without + # comparing it against the source. + evil_src = tmp_path / f"evil_{source_path.stem}.py" + evil_src.write_text(payload) + py_compile.compile( + str(evil_src), + cfile=importlib.util.cache_from_source(str(source_path)), + dfile=str(source_path), + invalidation_mode=py_compile.PycInvalidationMode.UNCHECKED_HASH, + ) + + +def test_import_ignores_bytecode(tmp_path): + variant_dir = _write_variant(tmp_path) + (variant_dir / "__init__.py").write_text("from ._sub import WHO as SUB_WHO\nWHO = 'source'\n") + (variant_dir / "_sub.py").write_text("WHO = 'source'\n") + _plant_pyc(variant_dir / "__init__.py", "from ._sub import WHO as SUB_WHO\nWHO = 'pyc'\n", tmp_path) + _plant_pyc(variant_dir / "_sub.py", "WHO = 'pyc'\n", tmp_path) + # Sourceless bytecode must not be importable either. + py_compile.compile(str(variant_dir / "_sub.py"), cfile=str(variant_dir / "_payload.pyc")) + + _loaded_kernels.pop(variant_dir, None) + try: + module = _import_from_path(variant_dir, deps={}) + assert module.WHO == "source" + assert module.SUB_WHO == "source" + with pytest.raises(ImportError): + importlib.import_module("broken_1_cuda._payload") + finally: + _loaded_kernels.pop(variant_dir, None) + for name in [m for m in sys.modules if m.startswith("broken_1_cuda")]: + sys.modules.pop(name)