File
Blob: src/pyodide/internal/topLevelEntropy/import_patch_manager.py
| 1 | """ |
| 2 | A metapath finder which calls get_import_context(module_name). If it returns a |
| 3 | value that is not None, this is interpreted as a context manager that should be |
| 4 | used when executing the module top level scope. |
| 5 | |
| 6 | When we're done, we put back the original module. The wrapper module and wrapper |
| 7 | stubs will persist in the wild, so we need to make sure they behave the same way |
| 8 | as the originals after we put them back. This is controlled by the |
| 9 | IN_REQUEST_CONTEXT variable. |
| 10 | """ |
| 11 | |
| 12 | import sys |
| 13 | from collections.abc import Callable |
| 14 | from contextlib import AbstractContextManager, nullcontext |
| 15 | from dataclasses import dataclass |
| 16 | from functools import partial, wraps |
| 17 | from typing import TYPE_CHECKING |
| 18 | |
| 19 | if TYPE_CHECKING: |
| 20 | from importlib.abc import Loader |
| 21 | from importlib.machinery import ModuleSpec |
| 22 | from types import ModuleType |
| 23 | |
| 24 | CreateImportContext = Callable[[ModuleSpec], AbstractContextManager] |
| 25 | ExecImportContext = Callable[[ModuleType], AbstractContextManager] |
| 26 | Handler = Callable[[ModuleType], None] |
| 27 | else: |
| 28 | from typing import Any |
| 29 | |
| 30 | Handler = Any |
| 31 | CreateImportContext = Any |
| 32 | ExecImportContext = Any |
| 33 | Loader = Any |
| 34 | ModuleSpec = Any |
| 35 | ModuleType = Any |
| 36 | |
| 37 | |
| 38 | @dataclass |
| 39 | class PatchInfo: |
| 40 | create: CreateImportContext = nullcontext |
| 41 | exec: ExecImportContext = nullcontext |
| 42 | after_snapshot: Handler | None = None |
| 43 | before_first_request: Handler | None = None |
| 44 | |
| 45 | |
| 46 | patches: dict[str, PatchInfo] = {} |
| 47 | |
| 48 | |
| 49 | # The public-facing interface here is: |
| 50 | # * register_create_patch |
| 51 | # * register_exec_patch |
| 52 | # * register_before_first_request |
| 53 | # * block_calls |
| 54 | |
| 55 | |
| 56 | def register_create_patch( |
| 57 | name: str, context_manager: CreateImportContext | None = None |
| 58 | ) -> None: |
| 59 | """This registers a context_manager that will be used around the create_module call when the |
| 60 | package named "name" is imported. |
| 61 | |
| 62 | It can either be used like register_exec_patch("cryptography.exceptions", rust_package_context) |
| 63 | or as a decorator. |
| 64 | """ |
| 65 | if context_manager is None: |
| 66 | return partial(register_create_patch, name) |
| 67 | d = patches.setdefault(name, PatchInfo()) |
| 68 | d.create = context_manager |
| 69 | return context_manager |
| 70 | |
| 71 | |
| 72 | def register_exec_patch( |
| 73 | name: str, context_manager: ExecImportContext | None = None |
| 74 | ) -> None: |
| 75 | """This registers a context_manager that will be used around the exec_module call when the |
| 76 | package named "name" is imported. |
| 77 | |
| 78 | It can either be used like register_create_patch("tiktoken._tiktoken", rust_package_context) |
| 79 | or as a decorator. |
| 80 | """ |
| 81 | if context_manager is None: |
| 82 | return partial(register_exec_patch, name) |
| 83 | d = patches.setdefault(name, PatchInfo()) |
| 84 | d.exec = context_manager |
| 85 | return context_manager |
| 86 | |
| 87 | |
| 88 | # Question: How do I decide whether to use register_create_patch or register_exec_patch? |
| 89 | # |
| 90 | # Answer: For pure Python packages, use register_exec_patch(). For extension modules, figure it out |
| 91 | # by trial and error. |
| 92 | |
| 93 | |
| 94 | def register_after_snapshot(name: str, handler: Handler | None = None): |
| 95 | if handler is None: |
| 96 | return partial(register_after_snapshot, name) |
| 97 | d = patches.setdefault(name, PatchInfo()) |
| 98 | d.after_snapshot = handler |
| 99 | return handler |
| 100 | |
| 101 | |
| 102 | def register_before_first_request(name: str, handler: Handler | None = None) -> None: |
| 103 | """This registers a callback that will be called before the first request if the package named |
| 104 | "name" was imported. Used for reseeding rng for instance. |
| 105 | """ |
| 106 | if handler is None: |
| 107 | return partial(register_before_first_request, name) |
| 108 | d = patches.setdefault(name, PatchInfo()) |
| 109 | d.before_first_request = handler |
| 110 | return handler |
| 111 | |
| 112 | |
| 113 | after_snapshot_handlers: list[Handler] = [] |
| 114 | before_first_request_handlers: list[Handler] = [] |
| 115 | |
| 116 | |
| 117 | class PatchLoader: |
| 118 | """Loader that calls the original exec_module in the given context manager""" |
| 119 | |
| 120 | def __init__(self, orig_loader: Loader, patch_info: PatchInfo): |
| 121 | self.orig_loader = orig_loader |
| 122 | self.patch_info = patch_info |
| 123 | |
| 124 | def __getattr__(self, name): |
| 125 | return getattr(self.orig_loader, name) |
| 126 | |
| 127 | def create_module(self, spec: ModuleSpec) -> ModuleType | None: |
| 128 | with self.patch_info.create(spec): |
| 129 | return self.orig_loader.create_module(spec) |
| 130 | |
| 131 | def exec_module(self, module: ModuleType) -> None: |
| 132 | if handler := self.patch_info.after_snapshot: |
| 133 | after_snapshot_handlers.append(partial(handler, module)) |
| 134 | if handler := self.patch_info.before_first_request: |
| 135 | before_first_request_handlers.append(partial(handler, module)) |
| 136 | with self.patch_info.exec(module): |
| 137 | self.orig_loader.exec_module(module) |
| 138 | |
| 139 | |
| 140 | class PatchFinder: |
| 141 | """Finder that returns our PatchLoader if get_import_context returns an import |
| 142 | context for the module. Otherwise, return None. |
| 143 | """ |
| 144 | |
| 145 | def invalidate_caches(self): |
| 146 | pass |
| 147 | |
| 148 | def find_spec( |
| 149 | self, |
| 150 | fullname: str, |
| 151 | path, |
| 152 | target, |
| 153 | ): |
| 154 | import_context = patches.get(fullname, None) |
| 155 | if import_context is None: |
| 156 | # Not ours |
| 157 | return None |
| 158 | |
| 159 | for finder in sys.meta_path: |
| 160 | if isinstance(finder, PatchFinder): |
| 161 | # Avoid infinite recurse. Presumably this is the first entry. |
| 162 | continue |
| 163 | spec = finder.find_spec(fullname, path, target) |
| 164 | if spec: |
| 165 | # Found original module spec |
| 166 | break |
| 167 | else: |
| 168 | # Not found. This is going to be an ImportError. |
| 169 | return None |
| 170 | # Overwrite the loader with our wrapped loader |
| 171 | spec.loader = PatchLoader(spec.loader, import_context) |
| 172 | return spec |
| 173 | |
| 174 | @staticmethod |
| 175 | def install(): |
| 176 | sys.meta_path.insert(0, PatchFinder()) |
| 177 | |
| 178 | @staticmethod |
| 179 | def remove(): |
| 180 | for idx, val in enumerate(sys.meta_path): # noqa:B007 |
| 181 | if isinstance(val, PatchFinder): |
| 182 | break |
| 183 | del sys.meta_path[idx] |
| 184 | |
| 185 | |
| 186 | def install_import_patch_manager(): |
| 187 | PatchFinder.install() |
| 188 | |
| 189 | |
| 190 | def remove_import_patch_manager(): |
| 191 | PatchFinder.remove() |
| 192 | unblock_calls() |
| 193 | |
| 194 | |
| 195 | # We remove the metapath entry and replace the patched sys.modules entries with |
| 196 | # the original modules before the request context, but the patched copies can |
| 197 | # still be used from top level imports. When IN_REQUEST_CONTEXT is True, we need |
| 198 | # to make sure that our patches behave like the original imports. |
| 199 | IN_REQUEST_CONTEXT = False |
| 200 | # Keep track of the unblocked modules so we can put them backk into sys.modules |
| 201 | # when we're done. |
| 202 | ORIG_MODULES = {} |
| 203 | |
| 204 | |
| 205 | def block_calls(module, *, allowlist=()): |
| 206 | """Make top level calls to methods from the module that are not in allowlist fail. |
| 207 | |
| 208 | It gets removed automatically before the first request. Generally used with |
| 209 | register_before_first_request. |
| 210 | """ |
| 211 | sys.modules[module.__name__] = BlockedCallModule(module, allowlist) |
| 212 | ORIG_MODULES[module.__name__] = module |
| 213 | |
| 214 | |
| 215 | def unblock_calls(): |
| 216 | # Remove the patches when we're ready to enable entropy calls. |
| 217 | global IN_REQUEST_CONTEXT |
| 218 | |
| 219 | IN_REQUEST_CONTEXT = True |
| 220 | for name, val in ORIG_MODULES.items(): |
| 221 | sys.modules[name] = val |
| 222 | |
| 223 | |
| 224 | class BlockedCallModule: |
| 225 | """A proxy class that wraps a module that we want to block calls to |
| 226 | |
| 227 | Attribute access is passed on to the original module but if the result is a |
| 228 | callable that isn't in the allow list, we wrap it with a function that |
| 229 | raises an error unless IN_REQUEST_CONTEXT is true. |
| 230 | |
| 231 | Note that because we define __getattribute__ and __setattr__, we cannot do |
| 232 | direct reads or assignments e.g., `self.a = 1`. This risks recursion errors |
| 233 | if there is a typo. Instead, we have to call super().__setattr__. |
| 234 | |
| 235 | This has the advantage that it avoids name clashes if the proxied module |
| 236 | actually defines variables called _mod or _allow_list. |
| 237 | """ |
| 238 | |
| 239 | def __init__(self, module, allowlist): |
| 240 | super().__setattr__("_mod", module) |
| 241 | super().__setattr__("_allow_list", allowlist) |
| 242 | |
| 243 | def __getattribute__(self, key): |
| 244 | mod = super().__getattribute__("_mod") |
| 245 | orig = getattr(mod, key) |
| 246 | if IN_REQUEST_CONTEXT: |
| 247 | return orig |
| 248 | if not callable(orig): |
| 249 | return orig |
| 250 | |
| 251 | if key in super().__getattribute__("_allow_list"): |
| 252 | return orig |
| 253 | |
| 254 | # If we aren't in a request scope, the value is a callable, and it's not |
| 255 | # in the allow_list, return a wrapper that raises an error if it's |
| 256 | # called before entering the request scope. |
| 257 | # TODO: this doesn't wrap classes correctly, does it matter? |
| 258 | @wraps(orig) |
| 259 | def wrapper(*args, **kwargs): |
| 260 | if not IN_REQUEST_CONTEXT: |
| 261 | raise RuntimeError( |
| 262 | f"Cannot use {mod.__name__}.{key}() outside of request context" |
| 263 | ) |
| 264 | return orig(*args, **kwargs) |
| 265 | |
| 266 | return wrapper |
| 267 | |
| 268 | def __setattr__(self, key, val): |
| 269 | mod = super().__getattribute__("_mod") |
| 270 | setattr(mod, key, val) |
| 271 | |
| 272 | def __dir__(self): |
| 273 | mod = super().__getattribute__("_mod") |
| 274 | return dir(mod) |