From a439aa0472bfd5ab2ed9eb226a2db69137010855 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 31 Aug 2026 00:55:34 +0800 Subject: [PATCH 1/2] Add regression tests for sticky transition validation --- ...rete_state_sticky_transition_validation.py | 38 +++++++++++++++++++ 1 file changed, 38 insertions(+) create mode 100644 tests/filters/test_discrete_state_sticky_transition_validation.py diff --git a/tests/filters/test_discrete_state_sticky_transition_validation.py b/tests/filters/test_discrete_state_sticky_transition_validation.py new file mode 100644 index 000000000..07e115b9d --- /dev/null +++ b/tests/filters/test_discrete_state_sticky_transition_validation.py @@ -0,0 +1,38 @@ +import numpy as np +import pytest + +from pyrecest.filters.discrete_state import sticky_mode_transition_matrix + + +@pytest.mark.parametrize( + "stickiness", + [ + True, + False, + "0.5", + 0.5 + 0j, + np.array([0.5]), + np.array([0.2, 0.8]), + np.nan, + np.inf, + -np.inf, + -0.1, + 1.1, + ], +) +def test_sticky_mode_transition_matrix_rejects_invalid_stickiness(stickiness): + with pytest.raises(ValueError, match="stickiness"): + sticky_mode_transition_matrix(3, stickiness) + + +@pytest.mark.parametrize( + "stickiness", + [0.0, 1.0, 0.25, np.float64(0.75), np.array(0.5)], +) +def test_sticky_mode_transition_matrix_accepts_real_scalar_stickiness(stickiness): + parsed = float(np.asarray(stickiness).item()) + result = sticky_mode_transition_matrix(3, stickiness) + + expected = np.full((3, 3), (1.0 - parsed) / 2.0) + np.fill_diagonal(expected, parsed) + np.testing.assert_allclose(result, expected) From 19323b04a383c154c2ce570e697b4d9789a4ba36 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Mon, 31 Aug 2026 00:56:32 +0800 Subject: [PATCH 2/2] Validate sticky transition probability inputs --- .../filters/discrete_state/__init__.py | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/src/pyrecest/filters/discrete_state/__init__.py b/src/pyrecest/filters/discrete_state/__init__.py index c70efe33a..4faf364e7 100644 --- a/src/pyrecest/filters/discrete_state/__init__.py +++ b/src/pyrecest/filters/discrete_state/__init__.py @@ -249,6 +249,25 @@ def _validated_positive_integer(value: Any, name: str) -> int: return parsed +def _validated_unit_interval_scalar(value: Any, name: str) -> float: + try: + raw_value = np.asarray(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"{name} must be a finite real scalar in [0, 1]") from exc + if raw_value.shape != () or raw_value.dtype.kind in _REJECTED_STATE_KINDS: + raise ValueError(f"{name} must be a finite real scalar in [0, 1]") + scalar = raw_value.item() + if isinstance(scalar, _TEXT_TYPES + _BOOLEAN_TYPES + _COMPLEX_TYPES): + raise ValueError(f"{name} must be a finite real scalar in [0, 1]") + try: + parsed = float(scalar) + except (TypeError, ValueError, OverflowError) as exc: + raise ValueError(f"{name} must be a finite real scalar in [0, 1]") from exc + if not np.isfinite(parsed) or not 0.0 <= parsed <= 1.0: + raise ValueError(f"{name} must be a finite real scalar in [0, 1]") + return parsed + + def _validated_probability_vector( probabilities: Any, n_entries: int, @@ -301,6 +320,7 @@ def _validated_probability_vector( @wraps(_original_sticky_mode_transition_matrix) def sticky_mode_transition_matrix(n_modes, stickiness): n_modes = _validated_positive_integer(n_modes, "n_modes") + stickiness = _validated_unit_interval_scalar(stickiness, "stickiness") return _original_sticky_mode_transition_matrix(n_modes, stickiness)