diff --git a/jaxtyping/_import_hook.py b/jaxtyping/_import_hook.py index 0aec718..492b336 100644 --- a/jaxtyping/_import_hook.py +++ b/jaxtyping/_import_hook.py @@ -76,8 +76,10 @@ def _optimized_cache_from_source(typechecker_hash, /, path, debug_override=None) # Version 5: Added support for string-based `typechecker` argument. # Version 6: optimization tag now depends on `typechecker` argument, so that # changing the typechecker will hit a different cache. + # Version 7: Using the same md5 hash of the `typechecker` argument + # for importlib and decorator lookup. return cache_from_source( - path, debug_override, optimization=f"jaxtyping6{typechecker_hash}" + path, debug_override, optimization=f"jaxtyping7{typechecker_hash}" ) @@ -88,20 +90,63 @@ def _dot_lookup(*elements): return out -def _str_lookup(string): - module = ast.parse(string) - (expr,) = module.body - return expr.value +class Typechecker: + lookup = {} + + def __init__(self, typechecker): + self.ast = None + + if isinstance(typechecker, str): + # If the typechecker is a string, then we parse it + string_to_eval = ( + "def f(x, *args, **kwargs):\n" + + f" import {typechecker.split('.', 1)[0]}\n" + + f" return {typechecker}(x, *args, **kwargs)" + ) + + # md5 hashing instead of __hash__ + # because __hash__ is different for each Python session + self.hash = hashlib.md5(typechecker.encode("utf-8")).hexdigest() + + vars = {} + exec(string_to_eval, {}, vars) + Typechecker.lookup[self.hash] = vars["f"] + + elif typechecker is None: + # If it is None, ignore it silently (use dummy decorator) + self.hash = 0 + Typechecker.lookup[self.hash] = lambda x, *_, **__: x + else: + # Passed typechecker is invalid + raise TypeError( + "Jaxtyping typechecker has to be either a string or a None." + ) + + def get_hash(self): + return self.hash + + def get_ast(self): + # we compile AST only if we missed importlib cache + if self.ast is None: + self.ast = ( + ast.parse( + f"@jaxtyping._import_hook.Typechecker.lookup['{self.hash}']\n" + "def _():\n ..." + ) + .body[0] + .decorator_list[0] + ) + + return self.ast -class _JaxtypingTransformer(ast.NodeVisitor): - def __init__(self, *, typechecker) -> None: +class JaxtypingTransformer(ast.NodeVisitor): + def __init__(self, *, typechecker: Typechecker) -> None: self._parents: list[ast.AST] = [] self._typechecker = typechecker def visit_Module(self, node: ast.Module): - # Insert "import typeguard; import jaxtping" after any "from __future__ ..." - # imports + # Insert "import jaxtyping" after any "from __future__ ..." imports for i, child in enumerate(node.body): if isinstance(child, ast.ImportFrom) and child.module == "__future__": continue @@ -109,11 +154,6 @@ class _JaxtypingTransformer(ast.NodeVisitor): continue # module docstring else: node.body.insert(i, ast.Import(names=[ast.alias("jaxtyping", None)])) - if self._typechecker is not None: - typechecker_module, _ = self._typechecker.split(".", 1) - node.body.insert( - i, ast.Import(names=[ast.alias(typechecker_module, None)]) - ) break self._parents.append(node) @@ -123,11 +163,9 @@ class _JaxtypingTransformer(ast.NodeVisitor): def visit_ClassDef(self, node: ast.ClassDef): func = _dot_lookup("jaxtyping", "_decorator", "_jaxtyped_typechecker") - if self._typechecker is None: - args = [ast.Constant(None)] - else: - args = [_str_lookup(self._typechecker)] - node.decorator_list.insert(0, ast.Call(func, args, keywords=[])) + node.decorator_list.insert( + 0, ast.Call(func, [self._typechecker.get_ast()], keywords=[]) + ) self._parents.append(node) self.generic_visit(node) self._parents.pop() @@ -151,11 +189,11 @@ class _JaxtypingTransformer(ast.NodeVisitor): # FWIW, typeguard also wants to be at the end of the decorator list, as it # works by recompiling the wrapped function. node.decorator_list.append(_dot_lookup("jaxtyping", "jaxtyped")) - if self._typechecker is not None: - # Place at the end of the decorator list, as decorators - # frequently remove annotations from functions and we'd like to - # use those annotations. - node.decorator_list.append(_str_lookup(self._typechecker)) + # Place typechecker at the end of the decorator list, as decorators + # frequently remove annotations from functions and we'd like to + # use those annotations. + node.decorator_list.append(self._typechecker.get_ast()) + self._parents.append(node) self.generic_visit(node) self._parents.pop() @@ -163,12 +201,9 @@ class _JaxtypingTransformer(ast.NodeVisitor): class _JaxtypingLoader(SourceFileLoader): - def __init__(self, *args, typechecker, **kwargs): + def __init__(self, *args, typechecker: Typechecker, **kwargs): super().__init__(*args, **kwargs) self._typechecker = typechecker - self._typechecker_hash = hashlib.md5( - self._typechecker.encode("utf-8") - ).hexdigest() def source_to_code(self, data, path, *, _optimize=-1): source = decode_source(data) @@ -181,7 +216,7 @@ class _JaxtypingLoader(SourceFileLoader): dont_inherit=True, optimize=_optimize, ) - tree = _JaxtypingTransformer(typechecker=self._typechecker).visit(tree) + tree = JaxtypingTransformer(typechecker=self._typechecker).visit(tree) ast.fix_missing_locations(tree) return _call_with_frames_removed( compile, tree, path, "exec", dont_inherit=True, optimize=_optimize @@ -192,7 +227,7 @@ class _JaxtypingLoader(SourceFileLoader): # patch safe with patch( "importlib._bootstrap_external.cache_from_source", - ft.partial(_optimized_cache_from_source, self._typechecker_hash), + ft.partial(_optimized_cache_from_source, self._typechecker.get_hash()), ): return super().exec_module(module) @@ -204,7 +239,7 @@ class _JaxtypingFinder(MetaPathFinder): Should not be used directly, but rather via `install_import_hook`. """ - def __init__(self, modules, original_pathfinder, typechecker): + def __init__(self, modules, original_pathfinder, typechecker: Typechecker): self.modules = modules self._original_pathfinder = original_pathfinder self._typechecker = typechecker @@ -385,6 +420,7 @@ def install_import_hook(modules: Union[str, Sequence[str]], typechecker: Optiona else: raise RuntimeError("Cannot find a PathFinder in sys.meta_path") - hook = _JaxtypingFinder(modules, finder, typechecker) + wrapped_typechecker = Typechecker(typechecker) + hook = _JaxtypingFinder(modules, finder, wrapped_typechecker) sys.meta_path.insert(0, hook) return ImportHookManager(hook) diff --git a/jaxtyping/_ipython_extension.py b/jaxtyping/_ipython_extension.py index 58095ba..7960b2f 100644 --- a/jaxtyping/_ipython_extension.py +++ b/jaxtyping/_ipython_extension.py @@ -17,7 +17,7 @@ # IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN # CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. -from ._import_hook import _JaxtypingTransformer +from ._import_hook import JaxtypingTransformer, Typechecker try: @@ -27,17 +27,17 @@ try: class ChooseTypecheckerMagics(Magics): @line_magic("jaxtyping.typechecker") def typechecker(self, typechecker): - # remove old _JaxtypingTransformer, if present + # remove old JaxtypingTransformer, if present self.shell.ast_transformers = list( filter( - lambda x: not isinstance(x, _JaxtypingTransformer), + lambda x: not isinstance(x, JaxtypingTransformer), self.shell.ast_transformers, ) ) # add new one self.shell.ast_transformers.append( - _JaxtypingTransformer(typechecker=typechecker) + JaxtypingTransformer(typechecker=Typechecker(typechecker)) ) except Exception: diff --git a/test/conftest.py b/test/conftest.py index c69a0e3..ca4e5d4 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -48,3 +48,13 @@ def getkey(): return jr.PRNGKey(random.randint(0, 2**31 - 1)) return _getkey + + +@pytest.fixture(scope="module") +def beartype_or_skip(): + yield pytest.importorskip("beartype") + + +@pytest.fixture(scope="module") +def typeguard_or_skip(): + yield pytest.importorskip("typeguard") diff --git a/test/test_import_hook.py b/test/test_import_hook.py index ebe1c84..8ad2527 100644 --- a/test/test_import_hook.py +++ b/test/test_import_hook.py @@ -69,45 +69,35 @@ def _test_import_hook(importhook_tempdir, typechecker): counter = jaxtyping._test_import_hook_counter stem = f"import_hook_tester{counter}" shutil.copyfile(_here / "import_hook_tester.py", importhook_tempdir / f"{stem}.py") + + importlib.invalidate_caches() with jaxtyping.install_import_hook(stem, typechecker): importlib.import_module(stem) assert counter + 1 == jaxtyping._test_import_hook_counter -def test_import_hook_typeguard_old(importhook_tempdir): - _test_import_hook(importhook_tempdir, ("typeguard", "typechecked")) +# Tests start below... -def test_import_hook_typeguard(importhook_tempdir): +def test_import_hook_typeguard(importhook_tempdir, typeguard_or_skip): _test_import_hook(importhook_tempdir, "typeguard.typechecked") -def test_import_hook_beartype_old(importhook_tempdir): - try: - import beartype # noqa: F401 - except ImportError: - pytest.skip("Beartype not installed") - else: - _test_import_hook(importhook_tempdir, ("beartype", "beartype")) +def test_import_hook_beartype(importhook_tempdir, beartype_or_skip): + _test_import_hook(importhook_tempdir, "beartype.beartype") -def test_import_hook_beartype(importhook_tempdir): - try: - import beartype # noqa: F401 - except ImportError: - pytest.skip("Beartype not installed") - else: - _test_import_hook(importhook_tempdir, "beartype.beartype") +def test_import_hook_beartype_full(importhook_tempdir, beartype_or_skip): + bearchecker = "beartype.beartype(conf=beartype.BeartypeConf(strategy=beartype.BeartypeStrategy.On))" # noqa: E501 + _test_import_hook(importhook_tempdir, bearchecker) -def test_import_hook_beartype_full(importhook_tempdir): - try: - import beartype # noqa: F401 - except ImportError: - pytest.skip("Beartype not installed") - else: - bearchecker = "beartype.beartype(conf=beartype.BeartypeConf(strategy=beartype.BeartypeStrategy.On))" # noqa: E501 - _test_import_hook(importhook_tempdir, bearchecker) +def test_import_hook_typeguard_old(importhook_tempdir, typeguard_or_skip): + _test_import_hook(importhook_tempdir, ("typeguard", "typechecked")) + + +def test_import_hook_beartype_old(importhook_tempdir, beartype_or_skip): + _test_import_hook(importhook_tempdir, ("beartype", "beartype")) def test_import_hook_broken_checker(importhook_tempdir): @@ -115,8 +105,7 @@ def test_import_hook_broken_checker(importhook_tempdir): _test_import_hook(importhook_tempdir, "jaxtyping.does_not_exist") -def test_import_hook_transitive(importhook_tempdir): - typechecker = "typeguard.typechecked" +def test_import_hook_transitive(importhook_tempdir, typeguard_or_skip): counter = jaxtyping._test_import_hook_counter transitive_name = "jaxtyping_transitive_test" transitive_dir = importhook_tempdir / transitive_name @@ -127,6 +116,6 @@ def test_import_hook_transitive(importhook_tempdir): f.flush() importlib.invalidate_caches() - with jaxtyping.install_import_hook(transitive_name, typechecker): + with jaxtyping.install_import_hook(transitive_name, "typeguard.typechecked"): importlib.import_module(transitive_name) assert counter + 1 == jaxtyping._test_import_hook_counter diff --git a/test/test_ipython_extension.py b/test/test_ipython_extension.py index a42330d..1c3c7d8 100644 --- a/test/test_ipython_extension.py +++ b/test/test_ipython_extension.py @@ -14,7 +14,7 @@ def ip(session_ip): session_ip.run_cell(raw_cell="import jaxtyping") session_ip.run_line_magic(magic_name="load_ext", line="jaxtyping") session_ip.run_line_magic( - magic_name="jaxtyping.typechecker", line="beartype.beartype" + magic_name="jaxtyping.typechecker", line="typeguard.typechecked" ) yield session_ip