diff --git a/src/pyrecest/distributions/abstract_manifold_specific_distribution.py b/src/pyrecest/distributions/abstract_manifold_specific_distribution.py index 4efc8e3d4..02622ea22 100644 --- a/src/pyrecest/distributions/abstract_manifold_specific_distribution.py +++ b/src/pyrecest/distributions/abstract_manifold_specific_distribution.py @@ -114,7 +114,7 @@ class AbstractManifoldSpecificDistribution(ABC): """ def __init__(self, dim: int): - self._dim = dim + self.dim = dim @abstractmethod def get_manifold_size(self) -> float: diff --git a/tests/distributions/test_manifold_constructor_dimension_validation.py b/tests/distributions/test_manifold_constructor_dimension_validation.py new file mode 100644 index 000000000..1f7c053c1 --- /dev/null +++ b/tests/distributions/test_manifold_constructor_dimension_validation.py @@ -0,0 +1,18 @@ +import unittest + +from pyrecest.distributions.hypertorus.custom_hypertoroidal_distribution import ( + CustomHypertoroidalDistribution, +) + + +class TestManifoldConstructorDimensionValidation(unittest.TestCase): + def test_custom_hypertoroidal_rejects_nonpositive_dimensions(self): + for dim in (0, -1): + with self.subTest(dim=dim), self.assertRaisesRegex( + ValueError, "dim must be a positive integer" + ): + CustomHypertoroidalDistribution(lambda _xs: 1.0, dim) + + +if __name__ == "__main__": + unittest.main()