diff --git a/pyproject.toml b/pyproject.toml index 4cf4ba4..ee903c4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,6 +30,7 @@ classifiers = [ requires-python = ">=3.8" dependencies = [ "numpy>=1.0", + "packaging", "spglib>=1.14.1", ] license = {text = "The MIT license"} diff --git a/seekpath/__init__.py b/seekpath/__init__.py index 61a0410..48bd95d 100644 --- a/seekpath/__init__.py +++ b/seekpath/__init__.py @@ -45,7 +45,19 @@ class SupercellWarning(UserWarning): ) from .hpkot import EdgeCaseWarning, SymmetryDetectionError -from .brillouinzone import brillouinzone + + +def __getattr__(name): + """Import the Brillouin zone subpackage only when it is used. + + It needs the optional ``scipy`` dependency, so importing it eagerly would + make ``import seekpath`` fail without it. + """ + if name == 'brillouinzone': + import importlib + + return importlib.import_module(f'{__name__}.brillouinzone') + raise AttributeError(f'module {__name__!r} has no attribute {name!r}') __all__ = ( diff --git a/seekpath/brillouinzone/__init__.py b/seekpath/brillouinzone/__init__.py index e69de29..8b1ed3b 100644 --- a/seekpath/brillouinzone/__init__.py +++ b/seekpath/brillouinzone/__init__.py @@ -0,0 +1,9 @@ +"""Compute the Brillouin zone of a crystal. + +This subpackage needs the optional ``scipy`` dependency, installed with +``pip install seekpath[bz]``. +""" + +from .brillouinzone import BZ, get_BZ + +__all__ = ('BZ', 'get_BZ') diff --git a/seekpath/brillouinzone/brillouinzone.py b/seekpath/brillouinzone/brillouinzone.py index 8bd93b9..a984219 100644 --- a/seekpath/brillouinzone/brillouinzone.py +++ b/seekpath/brillouinzone/brillouinzone.py @@ -5,7 +5,14 @@ from typing import Union import numpy as np -from scipy.spatial import Voronoi, ConvexHull, Delaunay + +try: + from scipy.spatial import ConvexHull, Delaunay, Voronoi +except ImportError as exc: + raise ImportError( + 'The Brillouin zone module requires scipy, install it with ' + '`pip install seekpath[bz]`' + ) from exc def get_BZ( diff --git a/tests/test_optional_dependencies.py b/tests/test_optional_dependencies.py new file mode 100644 index 0000000..eeda574 --- /dev/null +++ b/tests/test_optional_dependencies.py @@ -0,0 +1,62 @@ +"""Test that seekpath works without its optional dependencies.""" + +import subprocess +import sys +import unittest + +# Make every import of scipy fail, as if it were not installed +BLOCK_SCIPY = "import sys; sys.modules['scipy'] = None; " + + +def run_python(code): + """Run ``code`` in a fresh interpreter, so no module is imported yet.""" + return subprocess.run( + [sys.executable, '-c', code], capture_output=True, text=True, check=False + ) + + +class TestWithoutScipy(unittest.TestCase): + """seekpath only needs scipy for the Brillouin zone subpackage.""" + + def test_get_path_without_scipy(self): + """``import seekpath`` and ``get_path`` must not need scipy.""" + result = run_python( + BLOCK_SCIPY + 'import seekpath; ' + 'print(seekpath.get_path(([[4.0, 0, 0], [0, 4.0, 0], [0, 0, 4.0]], ' + "[[0, 0, 0]], [1]))['bravais_lattice'])" + ) + self.assertEqual(result.returncode, 0, result.stderr) + self.assertEqual(result.stdout.strip(), 'cP') + + def test_brillouinzone_without_scipy(self): + """The Brillouin zone subpackage explains how to install scipy.""" + result = run_python(BLOCK_SCIPY + 'import seekpath; seekpath.brillouinzone') + self.assertNotEqual(result.returncode, 0) + self.assertIn('ImportError', result.stderr) + self.assertIn('pip install seekpath[bz]', result.stderr) + + +class TestBrillouinzoneAccess(unittest.TestCase): + """The Brillouin zone subpackage is reachable in all the supported ways.""" + + def test_access_paths_agree(self): + """``seekpath.brillouinzone.BZ`` is the same class however it is reached.""" + result = run_python( + 'import seekpath; ' + 'from seekpath.brillouinzone import brillouinzone; ' + 'from seekpath.brillouinzone import BZ; ' + 'print(seekpath.brillouinzone.BZ is brillouinzone.BZ is BZ)' + ) + self.assertEqual(result.returncode, 0, result.stderr) + self.assertEqual(result.stdout.strip(), 'True') + + def test_unknown_attribute(self): + """Other missing attributes still raise ``AttributeError``.""" + import seekpath + + with self.assertRaises(AttributeError): + seekpath.not_an_attribute # noqa: B018 + + +if __name__ == '__main__': + unittest.main()