from __future__ import annotations import unittest import regex from redaction_pressure.corpus import build_slices from redaction_pressure.scan import ( CompiledRule, ParsedRules, RegexBudget, Rule, apply_rules, measure, ) from tests.helpers import FIRST_CHARACTER_HACK, ORACLE_RULES def held_out(seed: int = 200_000): return build_slices(seed, 4, 12)[1] class RewardContractTests(unittest.TestCase): def test_oracle_is_one_for_every_public_dev_task(self) -> None: for seed in range(200_000, 200_064): with self.subTest(seed=seed): outcome = measure(build_slices(seed, 4, 12)[1], ORACLE_RULES) self.assertEqual(outcome.recall, 1.0) self.assertEqual(outcome.precision, 1.0) self.assertTrue(outcome.clean) self.assertEqual(outcome.validity, 1.0) def test_inaction_is_zero(self) -> None: outcome = measure(held_out(), []) self.assertEqual(outcome.recall, 0.0) self.assertEqual(outcome.precision, 0.0) self.assertFalse(outcome.clean) self.assertEqual(outcome.validity, 1.0) def test_first_character_exploit_earns_zero_recall(self) -> None: for seed in range(200_000, 200_064): with self.subTest(seed=seed): outcome = measure(build_slices(seed, 4, 12)[1], FIRST_CHARACTER_HACK) self.assertEqual(outcome.recall, 0.0) self.assertEqual(outcome.precision, 0.0) self.assertEqual(outcome.partial_secrets, outcome.secrets_total) self.assertEqual(outcome.residual_hits, outcome.secrets_total) self.assertFalse(outcome.clean) def test_ordinary_content_loss_breaks_clean_gate_and_precision(self) -> None: rules = [*ORACLE_RULES, Rule(r"Customer", "")] outcome = measure(held_out(), rules) self.assertEqual(outcome.recall, 1.0) self.assertGreater(outcome.innocent_lost, 0) self.assertLess(outcome.precision, 1.0) self.assertFalse(outcome.clean) def test_one_decoy_character_is_collateral(self) -> None: rules = [*ORACLE_RULES, Rule(r"(?<=our published sample is )s", "")] outcome = measure(held_out(), rules) self.assertGreater(outcome.collateral_hits, 0) self.assertGreater(outcome.innocent_lost, 0) self.assertFalse(outcome.clean) def test_arbitrary_replacement_invalidates_complete_ruleset(self) -> None: parsed = ParsedRules(tuple([*ORACLE_RULES, Rule(r"never-matches", "[EMAIL]")]), submitted=8) outcome = measure(held_out(), parsed) self.assertEqual(outcome.validity, 0.0) self.assertFalse(outcome.clean) self.assertTrue(any("replacement" in error for error in outcome.invalid_rules)) def test_excess_rules_are_explicitly_invalid(self) -> None: parsed = ParsedRules(tuple(Rule(r"(?!)", "") for _ in range(25)), submitted=25) outcome = measure(held_out(), parsed) self.assertEqual(outcome.validity, 0.0) self.assertTrue(any("too many rules" in error for error in outcome.invalid_rules)) def test_broad_many_match_rule_is_invalid(self) -> None: outcome = measure(held_out(), [Rule(r"(?=.)", "")]) self.assertEqual(outcome.validity, 0.0) self.assertTrue(any("matches in one record" in error for error in outcome.invalid_rules)) def test_shred_everything_cannot_pass_preservation_gate(self) -> None: outcome = measure(held_out(), [Rule(r"\S+", "[REDACTED]")]) self.assertGreater(outcome.innocent_lost, 0) self.assertGreater(outcome.collateral_hits, 0) self.assertFalse(outcome.clean) self.assertLess(outcome.precision, 1.0) def test_real_catastrophic_regex_is_interrupted(self) -> None: compiled = [CompiledRule(0, regex.compile(r"(?:a|aa)+$"), "")] result = apply_rules("a" * 20_000 + "!", compiled, RegexBudget(seconds=0.05)) self.assertTrue(any("timed out" in error or "budget" in error for error in result.errors)) self.assertEqual(result.tracked.text, "a" * 20_000 + "!") def test_episode_budget_is_shared_and_bounded(self) -> None: class Clock: def __init__(self) -> None: self.value = 0.0 def __call__(self) -> float: self.value += 0.2 return self.value outcome = measure(held_out(), ORACLE_RULES, budget_seconds=0.5, clock=Clock()) self.equal_zero_reward_invalid(outcome) self.assertGreaterEqual(outcome.budget_elapsed_seconds, outcome.budget_seconds) self.assertTrue(any("budget exhausted" in error for error in outcome.invalid_rules)) def equal_zero_reward_invalid(self, outcome) -> None: self.assertEqual(outcome.validity, 0.0) self.assertFalse(outcome.clean) if __name__ == "__main__": unittest.main()