Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 39 additions & 1 deletion kernels/src/kernels/importer.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import importlib
import importlib.machinery
import os
import sys
from dataclasses import dataclass
from pathlib import Path
Expand Down Expand Up @@ -62,6 +64,39 @@ def get_loaded_kernels() -> list[LoadedKernel]:
return list(_loaded_kernels.values())


class _SourceOnlyLoader(importlib.machinery.SourceFileLoader):

@danieldk danieldk Sep 24, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

But we don't consider local cache poisoning as an attack vector. Once the cache could be poisoned, the attacker could also just replace .py files after all, or the program loading the kernel, or LD_PRELOAD a library to intercept calls, etc.

As long as we don't download .pyc files we should be good.

"""
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],
Expand All @@ -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)
Expand Down
37 changes: 37 additions & 0 deletions kernels/tests/test_importer.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
import importlib
import json
import py_compile
import sys

import pytest
Expand Down Expand Up @@ -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)
Loading