diff --git a/test/import_hook_tester_beartype.py b/test/import_hook_tester.py similarity index 64% rename from test/import_hook_tester_beartype.py rename to test/import_hook_tester.py index b77b756..e9f4369 100644 --- a/test/import_hook_tester_beartype.py +++ b/test/import_hook_tester.py @@ -22,10 +22,13 @@ import dataclasses import equinox as eqx import jax.numpy as jnp import pytest +from helpers import ParamError, ReturnError +import jaxtyping from jaxtyping import Float32 -from .helpers import ParamError + +# Test that functions get checked def g(x: Float32[jnp.ndarray, " b"]): @@ -36,6 +39,8 @@ g(jnp.array([1.0])) with pytest.raises(ParamError): g(jnp.array(1)) +# Test that Equinox modules get checked + class M(eqx.Module): foo: int @@ -48,6 +53,8 @@ with pytest.raises(ParamError): with pytest.raises(ParamError): M(1, jnp.array(1.0)) +# Test that dataclasses get checked + @dataclasses.dataclass class D: @@ -60,3 +67,44 @@ with pytest.raises(ParamError): D(1.0, jnp.array([1.0])) with pytest.raises(ParamError): D(1, jnp.array(1.0)) + +# Test that methods get checked + + +class N(eqx.Module): + a: jnp.ndarray + + def __init__(self, foo: str): + self.a = jnp.array(1) + + def foo(self, x: jnp.ndarray): + pass + + def bar(self) -> jnp.ndarray: + return self.a + + +n = N("hi") +with pytest.raises(ParamError): + N(123) +with pytest.raises(ParamError): + n.foo("not_an_array_either") +bad_n = eqx.tree_at(lambda x: x.a, n, "not_an_array") +with pytest.raises(ReturnError): + bad_n.bar() + +# Test that converters work + + +class BadConverter(eqx.Module): + a: jnp.ndarray = eqx.field(converter=lambda x: x) + + +with pytest.raises(ParamError): + BadConverter(1) +with pytest.raises(ParamError): + BadConverter("asdf") + +# Record that we've finished our checks successfully + +jaxtyping._test_import_hook_counter += 1 diff --git a/test/import_hook_tester_beartype_full.py b/test/import_hook_tester_beartype_full.py deleted file mode 100644 index b77b756..0000000 --- a/test/import_hook_tester_beartype_full.py +++ /dev/null @@ -1,62 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - -import dataclasses - -import equinox as eqx -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from .helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -class M(eqx.Module): - foo: int - bar: Float32[jnp.ndarray, " a"] - - -M(1, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1, jnp.array(1.0)) - - -@dataclasses.dataclass -class D: - foo: int - bar: Float32[jnp.ndarray, " a"] - - -D(1, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1, jnp.array(1.0)) diff --git a/test/import_hook_tester_beartype_old.py b/test/import_hook_tester_beartype_old.py deleted file mode 100644 index b77b756..0000000 --- a/test/import_hook_tester_beartype_old.py +++ /dev/null @@ -1,62 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - -import dataclasses - -import equinox as eqx -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from .helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -class M(eqx.Module): - foo: int - bar: Float32[jnp.ndarray, " a"] - - -M(1, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1, jnp.array(1.0)) - - -@dataclasses.dataclass -class D: - foo: int - bar: Float32[jnp.ndarray, " a"] - - -D(1, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1, jnp.array(1.0)) diff --git a/test/import_hook_tester_broken_checker.py b/test/import_hook_tester_broken_checker.py deleted file mode 100644 index b77b756..0000000 --- a/test/import_hook_tester_broken_checker.py +++ /dev/null @@ -1,62 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - -import dataclasses - -import equinox as eqx -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from .helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -class M(eqx.Module): - foo: int - bar: Float32[jnp.ndarray, " a"] - - -M(1, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - M(1, jnp.array(1.0)) - - -@dataclasses.dataclass -class D: - foo: int - bar: Float32[jnp.ndarray, " a"] - - -D(1, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1.0, jnp.array([1.0])) -with pytest.raises(ParamError): - D(1, jnp.array(1.0)) diff --git a/test/import_hook_tester_transitive/__init__.py b/test/import_hook_tester_transitive/__init__.py deleted file mode 100644 index 3501c8a..0000000 --- a/test/import_hook_tester_transitive/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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 another_file # noqa: F401 diff --git a/test/import_hook_tester_transitive/another_file.py b/test/import_hook_tester_transitive/another_file.py deleted file mode 100644 index c4fc46d..0000000 --- a/test/import_hook_tester_transitive/another_file.py +++ /dev/null @@ -1,48 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from ..helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -# Typeguard 3.0 no longer supports this -# -# class M(eqx.Module): -# foo: int -# bar: Float32[jnp.ndarray, " a"] - - -# M(1, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1.0, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1, jnp.array(1.0)) diff --git a/test/import_hook_tester_typeguard.py b/test/import_hook_tester_typeguard.py deleted file mode 100644 index b184215..0000000 --- a/test/import_hook_tester_typeguard.py +++ /dev/null @@ -1,63 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - - -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from .helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -# Typeguard 3.0 no longer supports this. -# -# class M(eqx.Module): -# foo: int -# bar: Float32[jnp.ndarray, " a"] -# -# -# M(1, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1.0, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1, jnp.array(1.0)) -# -# -# -# @dataclasses.dataclass -# class D: -# foo: int -# bar: Float32[jnp.ndarray, " a"] -# -# -# D(1, jnp.array([1.0])) -# with pytest.raises(ParamError): -# D(1.0, jnp.array([1.0])) -# with pytest.raises(ParamError): -# D(1, jnp.array(1.0)) diff --git a/test/import_hook_tester_typeguard_old.py b/test/import_hook_tester_typeguard_old.py deleted file mode 100644 index b184215..0000000 --- a/test/import_hook_tester_typeguard_old.py +++ /dev/null @@ -1,63 +0,0 @@ -# Copyright (c) 2022 Google LLC -# -# Permission is hereby granted, free of charge, to any person obtaining a copy of -# this software and associated documentation files (the "Software"), to deal in -# the Software without restriction, including without limitation the rights to -# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of -# the Software, and to permit persons to whom the Software is furnished to do so, -# subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in all -# copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS -# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR -# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER -# 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. - - -import jax.numpy as jnp -import pytest - -from jaxtyping import Float32 - -from .helpers import ParamError - - -def g(x: Float32[jnp.ndarray, " b"]): - pass - - -g(jnp.array([1.0])) -with pytest.raises(ParamError): - g(jnp.array(1)) - - -# Typeguard 3.0 no longer supports this. -# -# class M(eqx.Module): -# foo: int -# bar: Float32[jnp.ndarray, " a"] -# -# -# M(1, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1.0, jnp.array([1.0])) -# with pytest.raises(ParamError): -# M(1, jnp.array(1.0)) -# -# -# -# @dataclasses.dataclass -# class D: -# foo: int -# bar: Float32[jnp.ndarray, " a"] -# -# -# D(1, jnp.array([1.0])) -# with pytest.raises(ParamError): -# D(1.0, jnp.array([1.0])) -# with pytest.raises(ParamError): -# D(1, jnp.array(1.0)) diff --git a/test/test_import_hook.py b/test/test_import_hook.py index 9f35fe4..6a1ef73 100644 --- a/test/test_import_hook.py +++ b/test/test_import_hook.py @@ -17,76 +17,114 @@ # 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. +import importlib +import importlib.metadata +import pathlib +import shutil +import sys +import tempfile + import pytest -from jaxtyping import install_import_hook +import jaxtyping -def test_import_hook_typeguard_old(): - hook = install_import_hook( - "test.import_hook_tester_typeguard_old", ("typeguard", "typechecked") +_here = pathlib.Path(__file__).resolve().parent + + +try: + typeguard_version = importlib.metadata.version("typeguard") +except Exception as e: + raise ImportError("Could not find typeguard version") from e +else: + try: + major, _, _ = typeguard_version.split(".") + major = int(major) + except Exception as e: + raise ImportError( + f"Unexpected typeguard version {typeguard_version}; not formatted as " + "`major.minor.patch`" + ) from e +if major != 2: + raise ImportError( + "jaxtyping's tests required typeguard version 2. (Versions 3 and 4 are both " + "known to have bugs.)" ) - with hook: - from . import import_hook_tester_typeguard_old # noqa: F401 -def test_import_hook_typeguard(): - hook = install_import_hook( - "test.import_hook_tester_typeguard", "typeguard.typechecked" - ) - with hook: - from . import import_hook_tester_typeguard # noqa: F401 +assert not hasattr(jaxtyping, "_test_import_hook_counter") +jaxtyping._test_import_hook_counter = 0 -def test_import_hook_beartype_old(): +@pytest.fixture(scope="module") +def importhook_tempdir(): + with tempfile.TemporaryDirectory() as dir: + sys.path.append(dir) + dir = pathlib.Path(dir) + shutil.copyfile(_here / "helpers.py", dir / "helpers.py") + yield dir + + +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") + 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")) + + +def test_import_hook_typeguard(importhook_tempdir): + _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: - hook = install_import_hook( - "test.import_hook_tester_beartype_old", ("beartype", "beartype") - ) - with hook: - from . import import_hook_tester_beartype_old # noqa: F401 + _test_import_hook(importhook_tempdir, ("beartype", "beartype")) -def test_import_hook_beartype(): +def test_import_hook_beartype(importhook_tempdir): try: import beartype # noqa: F401 except ImportError: pytest.skip("Beartype not installed") else: - hook = install_import_hook( - "test.import_hook_tester_beartype", "beartype.beartype" - ) - with hook: - from . import import_hook_tester_beartype # noqa: F401 + _test_import_hook(importhook_tempdir, "beartype.beartype") -def test_import_hook_beartype_full(): +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 - hook = install_import_hook("test.import_hook_tester_beartype_full", bearchecker) - with hook: - from . import import_hook_tester_beartype_full # noqa: F401 + _test_import_hook(importhook_tempdir, bearchecker) -def test_import_hook_transitive(): - hook = install_import_hook( - "test.import_hook_tester_transitive", "typeguard.typechecked" - ) - with hook: - from . import import_hook_tester_transitive # noqa: F401 +def test_import_hook_broken_checker(importhook_tempdir): + with pytest.raises(AttributeError): + _test_import_hook(importhook_tempdir, "jaxtyping.does_not_exist") -def test_import_hook_broken_checker(): - hook = install_import_hook( - "test.import_hook_tester_broken_checker", "jaxtyping.does_not_exist" - ) - with hook, pytest.raises(AttributeError): - from . import import_hook_tester_broken_checker # noqa: F401 +def test_import_hook_transitive(importhook_tempdir): + typechecker = "typeguard.typechecked" + counter = jaxtyping._test_import_hook_counter + transitive_name = "jaxtyping_transitive_test" + transitive_dir = importhook_tempdir / transitive_name + transitive_dir.mkdir() + shutil.copyfile(_here / "import_hook_tester.py", transitive_dir / "tester.py") + with open(transitive_dir / "__init__.py", "w") as f: + f.write("from . import tester") + f.flush() + with jaxtyping.install_import_hook(transitive_name, typechecker): + importlib.import_module(transitive_name) + assert counter + 1 == jaxtyping._test_import_hook_counter