mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-12 12:32:19 +08:00
Better handling of user-defined typechecker
This commit is contained in:
committed by
Patrick Kidger
parent
9c9635d4f3
commit
513a54b048
@@ -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")
|
||||
|
||||
+17
-28
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user