Files
OBLITERATUS/tests/test_numerical_contracts.py
T

64 lines
2.3 KiB
Python

"""Mutation-focused tests for pure numerical input contracts."""
from __future__ import annotations
import pytest
from obliteratus.analysis.numerical_contracts import (
validate_whitened_parameters,
validate_whitened_request,
)
@pytest.mark.parametrize("regularization_eps", [0, -1e-4, float("nan"), float("inf"), True, "1e-4"])
def test_regularization_must_be_a_finite_positive_number(regularization_eps):
with pytest.raises(ValueError, match="finite positive"):
validate_whitened_parameters(regularization_eps, 0.01)
@pytest.mark.parametrize(
"min_variance_ratio",
[-0.1, 1.0, float("nan"), float("inf"), True, "0.01"],
)
def test_variance_ratio_must_be_in_the_half_open_unit_interval(min_variance_ratio):
with pytest.raises(ValueError, match=r"interval \[0, 1\)"):
validate_whitened_parameters(1e-4, min_variance_ratio)
@pytest.mark.parametrize(
("regularization_eps", "min_variance_ratio", "expected"),
[(1, 0, (1.0, 0.0)), (1e-4, 0.999, (1e-4, 0.999))],
)
def test_valid_parameters_are_normalized_to_floats(
regularization_eps,
min_variance_ratio,
expected,
):
assert validate_whitened_parameters(regularization_eps, min_variance_ratio) == expected
@pytest.mark.parametrize(
("harmful_count", "harmless_count", "message"),
[(0, 0, "both be non-empty"), (1, 0, "both be non-empty"), (0, 1, "both be non-empty")],
)
def test_activation_sets_must_both_be_nonempty(harmful_count, harmless_count, message):
with pytest.raises(ValueError, match=message):
validate_whitened_request(harmful_count, harmless_count, 1)
@pytest.mark.parametrize(("harmful_count", "harmless_count"), [(1, 2), (3, 1)])
def test_activation_sets_must_have_equal_sample_counts(harmful_count, harmless_count):
with pytest.raises(ValueError, match=f"got {harmful_count} and {harmless_count}"):
validate_whitened_request(harmful_count, harmless_count, 1)
@pytest.mark.parametrize("n_directions", [-1, 0, 1.5, True, "1"])
def test_direction_count_must_be_a_positive_integer(n_directions):
with pytest.raises(ValueError, match="positive integer"):
validate_whitened_request(2, 2, n_directions)
@pytest.mark.parametrize("n_directions", [1, 2, 100])
def test_valid_direction_count_is_returned(n_directions):
assert validate_whitened_request(2, 2, n_directions) == n_directions