diff --git a/src/pyrecest/filters/velocity_aided_mem_qkf_tracker.py b/src/pyrecest/filters/velocity_aided_mem_qkf_tracker.py index b1c5a2bb9..1ce4480d0 100644 --- a/src/pyrecest/filters/velocity_aided_mem_qkf_tracker.py +++ b/src/pyrecest/filters/velocity_aided_mem_qkf_tracker.py @@ -1,5 +1,9 @@ from __future__ import annotations +from math import isfinite + +import numpy as np + # pylint: disable=no-name-in-module,no-member,too-many-arguments # pylint: disable=too-many-positional-arguments,too-many-locals from pyrecest.backend import arctan2, array, cos, maximum, sin @@ -7,6 +11,30 @@ from .mem_qkf_tracker import MEMQKFTracker +def _as_finite_float(value, name: str) -> float: + """Return a finite scalar float without accepting boolean controls.""" + value_array = np.asarray(value) + if value_array.shape != () or value_array.dtype == np.bool_: + raise ValueError(f"{name} must be a finite scalar") + scalar = value_array.item() + if isinstance(scalar, (bool, np.bool_)): + raise ValueError(f"{name} must be a finite scalar") + try: + parsed = float(scalar) + except (TypeError, ValueError, OverflowError) as exc: + raise ValueError(f"{name} must be a finite scalar") from exc + if not isfinite(parsed): + raise ValueError(f"{name} must be a finite scalar") + return parsed + + +def _as_bool_flag(value, name: str) -> bool: + """Return a strict Python boolean for public configuration flags.""" + if isinstance(value, (bool, np.bool_)): + return bool(value) + raise ValueError(f"{name} must be a boolean") + + class VelocityAidedMEMQKFTracker(MEMQKFTracker): """MEM-QKF with a soft axial velocity-heading pseudo-measurement. @@ -66,21 +94,36 @@ def __init__( super().__init__(*args, **kwargs) self.velocity_indices = self._normalize_velocity_indices(velocity_indices) - self.speed_threshold = float(speed_threshold) + self.speed_threshold = _as_finite_float(speed_threshold, "speed_threshold") if self.speed_threshold < 0.0: raise ValueError("speed_threshold must be non-negative") - self.orientation_offset = float(orientation_offset) - self.heading_noise_variance = float(heading_noise_variance) + self.orientation_offset = _as_finite_float( + orientation_offset, + "orientation_offset", + ) + self.heading_noise_variance = _as_finite_float( + heading_noise_variance, + "heading_noise_variance", + ) if self.heading_noise_variance < 0.0: raise ValueError("heading_noise_variance must be non-negative") - self.minimum_heading_variance = float(minimum_heading_variance) + self.minimum_heading_variance = _as_finite_float( + minimum_heading_variance, + "minimum_heading_variance", + ) if self.minimum_heading_variance <= 0.0: raise ValueError("minimum_heading_variance must be positive") - self.apply_heading_on_prediction = bool(apply_heading_on_prediction) - self.use_heading_constraint = bool(use_heading_constraint) + self.apply_heading_on_prediction = _as_bool_flag( + apply_heading_on_prediction, + "apply_heading_on_prediction", + ) + self.use_heading_constraint = _as_bool_flag( + use_heading_constraint, + "use_heading_constraint", + ) self._heading_update_pending = False def _normalize_velocity_indices(self, velocity_indices): @@ -211,7 +254,10 @@ def update( """Update and fuse the velocity-heading pseudo-measurement once per scan.""" old_use_heading_constraint = self.use_heading_constraint if use_heading_constraint is not None: - self.use_heading_constraint = bool(use_heading_constraint) + self.use_heading_constraint = _as_bool_flag( + use_heading_constraint, + "use_heading_constraint", + ) self._heading_update_pending = True try: super().update( diff --git a/tests/filters/test_velocity_aided_mem_qkf_validation.py b/tests/filters/test_velocity_aided_mem_qkf_validation.py new file mode 100644 index 000000000..332baab1f --- /dev/null +++ b/tests/filters/test_velocity_aided_mem_qkf_validation.py @@ -0,0 +1,66 @@ +# pylint: disable=protected-access +import numpy as np +import pytest +from pyrecest.backend import array, diag, eye +from pyrecest.filters.velocity_aided_mem_qkf_tracker import VelocityAidedMEMQKFTracker + + +def _make_tracker(**kwargs): + config = { + "kinematic_state": array([0.0, 0.0, 3.0, 4.0]), + "covariance": diag(array([1.0, 1.0, 0.04, 0.04])), + "shape_state": array([0.0, 2.0, 1.0]), + "shape_covariance": diag(array([0.5, 0.2, 0.2])), + "default_meas_noise_cov": 0.05 * eye(2), + "heading_noise_variance": 0.25, + } + config.update(kwargs) + return VelocityAidedMEMQKFTracker(**config) + + +@pytest.mark.parametrize( + "parameter", + ( + "speed_threshold", + "orientation_offset", + "heading_noise_variance", + "minimum_heading_variance", + ), +) +@pytest.mark.parametrize("value", (np.nan, np.inf, -np.inf)) +def test_heading_scalar_controls_reject_nonfinite_values(parameter, value): + with pytest.raises(ValueError, match=parameter): + _make_tracker(**{parameter: value}) + + +@pytest.mark.parametrize( + "parameter", + ("apply_heading_on_prediction", "use_heading_constraint"), +) +@pytest.mark.parametrize("value", ("False", 0, 1, None, np.array(False))) +def test_heading_boolean_controls_reject_non_boolean_values(parameter, value): + with pytest.raises(ValueError, match=parameter): + _make_tracker(**{parameter: value}) + + +def test_heading_boolean_controls_accept_numpy_booleans(): + tracker = _make_tracker( + apply_heading_on_prediction=np.bool_(False), + use_heading_constraint=np.bool_(True), + ) + + assert tracker.apply_heading_on_prediction is False + assert tracker.use_heading_constraint is True + + +def test_per_update_heading_override_rejects_non_boolean_without_mutating_state(): + tracker = _make_tracker() + + with pytest.raises(ValueError, match="use_heading_constraint"): + tracker.update( + array([[1.0, 0.2]]), + use_heading_constraint="False", + ) + + assert tracker.use_heading_constraint is True + assert tracker._heading_update_pending is False