117 lines
4.8 KiB
Python
117 lines
4.8 KiB
Python
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()
|