diff --git a/src/pyrecest/filters/ekf_spline_tracker.py b/src/pyrecest/filters/ekf_spline_tracker.py index b9ddd48de..b9d24e031 100644 --- a/src/pyrecest/filters/ekf_spline_tracker.py +++ b/src/pyrecest/filters/ekf_spline_tracker.py @@ -1,4 +1,5 @@ # pylint: disable=duplicate-code,no-member,no-name-in-module,too-many-lines +from math import isfinite from numbers import Integral from pyrecest.backend import ( @@ -125,8 +126,8 @@ def __init__( "closest_point_iterations", 0, ) - if self.finite_difference_step <= 0.0: - raise ValueError("finite_difference_step must be positive") + if not isfinite(self.finite_difference_step) or self.finite_difference_step <= 0.0: + raise ValueError("finite_difference_step must be finite and positive") self.last_quadratic_form = None self._sync_state_views() @@ -585,4 +586,4 @@ def get_bounding_box(self, n=100): } -EkfSplineTracker = EKFSplineTracker +EkfSplineTracker = EKFSplineTracker \ No newline at end of file diff --git a/src/pyrecest/filters/mem_soekf_tracker.py b/src/pyrecest/filters/mem_soekf_tracker.py index 0e69834d7..3ede7f99a 100644 --- a/src/pyrecest/filters/mem_soekf_tracker.py +++ b/src/pyrecest/filters/mem_soekf_tracker.py @@ -1,5 +1,7 @@ from __future__ import annotations +from math import isfinite + # pylint: disable=no-name-in-module,no-member,duplicate-code,too-many-locals from pyrecest.backend import ( array, @@ -29,8 +31,8 @@ class MEMSOEKFTracker(MEMEKFTracker): def __init__(self, *args, finite_difference_step=1e-5, **kwargs): super().__init__(*args, **kwargs) self.finite_difference_step = float(finite_difference_step) - if self.finite_difference_step <= 0.0: - raise ValueError("finite_difference_step must be positive") + if not isfinite(self.finite_difference_step) or self.finite_difference_step <= 0.0: + raise ValueError("finite_difference_step must be finite and positive") @staticmethod def _extent_transform_from_shape(shape_state): diff --git a/tests/filters/test_finite_difference_step_validation.py b/tests/filters/test_finite_difference_step_validation.py new file mode 100644 index 000000000..451d865d6 --- /dev/null +++ b/tests/filters/test_finite_difference_step_validation.py @@ -0,0 +1,48 @@ +import unittest + +# pylint: disable=no-name-in-module,no-member +import pyrecest.backend +from pyrecest.backend import array, diag +from pyrecest.filters import EKFSplineTracker, MEMSOEKFTracker + + +@unittest.skipIf( + pyrecest.backend.__backend_name__ != "numpy", + reason="Finite-difference tracker validation tests use the NumPy backend.", +) +class TestFiniteDifferenceStepValidation(unittest.TestCase): + @staticmethod + def _make_mem_soekf(finite_difference_step): + return MEMSOEKFTracker( + array([0.0, 0.0, 1.0, -1.0]), + diag(array([0.1, 0.1, 0.01, 0.01])), + array([0.0, 2.0, 1.0]), + diag(array([0.01, 0.1, 0.2])), + measurement_matrix=array( + [ + [1.0, 0.0, 0.0, 0.0], + [0.0, 1.0, 0.0, 0.0], + ] + ), + finite_difference_step=finite_difference_step, + ) + + def test_mem_soekf_rejects_nonfinite_step(self): + for invalid_step in (float("nan"), float("inf"), float("-inf")): + with self.subTest(step=invalid_step), self.assertRaisesRegex( + ValueError, + "finite_difference_step", + ): + self._make_mem_soekf(invalid_step) + + def test_ekf_spline_rejects_nonfinite_step(self): + for invalid_step in (float("nan"), float("inf"), float("-inf")): + with self.subTest(step=invalid_step), self.assertRaisesRegex( + ValueError, + "finite_difference_step", + ): + EKFSplineTracker(finite_difference_step=invalid_step) + + +if __name__ == "__main__": + unittest.main()