Skip to content
File

Blob: src/pyodide/internal/topLevelEntropy/import_patch_manager.py

python275 lines
1"""
2A metapath finder which calls get_import_context(module_name). If it returns a
3value that is not None, this is interpreted as a context manager that should be
4used when executing the module top level scope.
5
6When we're done, we put back the original module. The wrapper module and wrapper
7stubs will persist in the wild, so we need to make sure they behave the same way
8as the originals after we put them back. This is controlled by the
9IN_REQUEST_CONTEXT variable.
10"""
11 
12import sys
13from collections.abc import Callable
14from contextlib import AbstractContextManager, nullcontext
15from dataclasses import dataclass
16from functools import partial, wraps
17from typing import TYPE_CHECKING
18 
19if 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]
27else:
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
39class PatchInfo:
40 create: CreateImportContext = nullcontext
41 exec: ExecImportContext = nullcontext
42 after_snapshot: Handler | None = None
43 before_first_request: Handler | None = None
44 
45 
46patches: 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 
56def 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 
72def 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 
94def 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 
102def 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 
113after_snapshot_handlers: list[Handler] = []
114before_first_request_handlers: list[Handler] = []
115 
116 
117class 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 
140class 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 
186def install_import_patch_manager():
187 PatchFinder.install()
188 
189 
190def 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.
199IN_REQUEST_CONTEXT = False
200# Keep track of the unblocked modules so we can put them backk into sys.modules
201# when we're done.
202ORIG_MODULES = {}
203 
204 
205def 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 
215def 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 
224class 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)