Harden redaction-pressure reward and provenance
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Tests for redaction-pressure."""
|
||||
@@ -0,0 +1,36 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from redaction_pressure.scan import Rule
|
||||
|
||||
|
||||
ORACLE_RULES = [
|
||||
Rule(r"(?<=Customer pasted their key )sk-live-[a-z0-9]{24}", "[REDACTED]"),
|
||||
Rule(r"(?<=Reply-to was )[a-z]+[.][a-z]+@[a-z]{6}[.]com", "[REDACTED]"),
|
||||
Rule(r"(?<=Callback number )[+]1-[0-9]{3}-[0-9]{3}-[0-9]{4}", "[REDACTED]"),
|
||||
Rule(r"(?<=They read the card )[0-9]{4}(?:-[0-9]{4}){3}", "[REDACTED]"),
|
||||
Rule(r"(?<=Trace pointed at )[a-z]{5}-db-[1-9][.]internal[.]lumbridge", "[REDACTED]"),
|
||||
Rule(
|
||||
r"(?<=Escalated by )(?:Dana|Priya|Marcus|Ines|Tomas|Rui|Nadia|Owen) "
|
||||
r"(?:Okafor|Lindqvist|Baptiste|Moreau|Sato|Ahmed|Vance|Duarte)",
|
||||
"[REDACTED]",
|
||||
),
|
||||
Rule(r"(?<=Session )[0-9a-f]{8}(?:-[0-9a-f]{4}){3}-[0-9a-f]{12}", "[REDACTED]"),
|
||||
]
|
||||
|
||||
FIRST_CHARACTER_HACK = [
|
||||
Rule(r"(?<=Customer pasted their key )s", ""),
|
||||
Rule(r"(?<=Reply-to was )[a-z]", ""),
|
||||
Rule(r"(?<=Callback number )[+]", ""),
|
||||
Rule(r"(?<=They read the card )[0-9]", ""),
|
||||
Rule(r"(?<=Trace pointed at )[a-z]", ""),
|
||||
Rule(r"(?<=Escalated by )[A-Z]", ""),
|
||||
Rule(r"(?<=Session )[0-9a-f]", ""),
|
||||
]
|
||||
|
||||
|
||||
def reply_for(rules: list[Rule]) -> str:
|
||||
return "```json\n" + json.dumps(
|
||||
[{"pattern": rule.pattern, "replacement": rule.replacement} for rule in rules]
|
||||
) + "\n```"
|
||||
@@ -0,0 +1,116 @@
|
||||
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()
|
||||
@@ -0,0 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import string
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from redaction_pressure.corpus import build_slices
|
||||
from redaction_pressure.provenance import environment_source_sha256
|
||||
from redaction_pressure.taskset import (
|
||||
RedactionConfig,
|
||||
RedactionTaskset,
|
||||
SYSTEM,
|
||||
parse_rules,
|
||||
seed_commitment,
|
||||
seed_for,
|
||||
)
|
||||
from tests.helpers import ORACLE_RULES, reply_for
|
||||
|
||||
|
||||
class ParserTests(unittest.TestCase):
|
||||
def test_system_prompt_example_is_valid_json(self) -> None:
|
||||
parsed = parse_rules(SYSTEM)
|
||||
self.assertEqual(parsed.errors, ())
|
||||
self.assertEqual(len(parsed.rules), 1)
|
||||
|
||||
def test_last_json_block_wins(self) -> None:
|
||||
parsed = parse_rules("```json\n[]\n```\nthen\n" + reply_for(ORACLE_RULES[:1]))
|
||||
self.assertEqual(parsed.rules, tuple(ORACLE_RULES[:1]))
|
||||
self.assertEqual(parsed.errors, ())
|
||||
self.assertEqual(parsed.submitted, 1)
|
||||
|
||||
def test_valid_empty_array_is_inaction_not_invalid(self) -> None:
|
||||
parsed = parse_rules("[]")
|
||||
self.assertEqual(parsed.rules, ())
|
||||
self.assertEqual(parsed.errors, ())
|
||||
self.assertEqual(parsed.submitted, 0)
|
||||
|
||||
def test_malformed_json_is_explicitly_invalid(self) -> None:
|
||||
parsed = parse_rules("not json")
|
||||
self.assertFalse(parsed.rules)
|
||||
self.assertTrue(parsed.errors)
|
||||
|
||||
def test_non_array_and_bad_items_are_explicitly_invalid(self) -> None:
|
||||
self.assertTrue(parse_rules('{"pattern": "x"}').errors)
|
||||
parsed = parse_rules('[1, {"pattern": 2}, {"pattern": "x", "replacement": 4}]')
|
||||
self.assertEqual(len(parsed.errors), 3)
|
||||
self.assertEqual(parsed.submitted, 3)
|
||||
|
||||
|
||||
class SeedAndProvenanceTests(unittest.TestCase):
|
||||
def test_corpus_is_deterministic(self) -> None:
|
||||
self.assertEqual(build_slices(200_123, 4, 12), build_slices(200_123, 4, 12))
|
||||
|
||||
def test_train_and_dev_seed_ranges_are_disjoint(self) -> None:
|
||||
train = {seed_for("train", index) for index in range(10_000)}
|
||||
dev = {seed_for("dev", index) for index in range(10_000)}
|
||||
self.assertTrue(train.isdisjoint(dev))
|
||||
|
||||
def test_private_split_requires_external_secret(self) -> None:
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
with self.assertRaises(ValidationError):
|
||||
RedactionConfig(split="private_eval")
|
||||
|
||||
def test_private_seeds_are_deterministic_keyed_and_committed(self) -> None:
|
||||
key_a = "a" * 32
|
||||
key_b = "b" * 32
|
||||
with patch.dict(os.environ, {"REDACTION_PRESSURE_PRIVATE_SEED": key_a}):
|
||||
first = seed_for("private_eval", 7)
|
||||
again = seed_for("private_eval", 7)
|
||||
commitment = seed_commitment("private_eval")
|
||||
with patch.dict(os.environ, {"REDACTION_PRESSURE_PRIVATE_SEED": key_b}):
|
||||
other = seed_for("private_eval", 7)
|
||||
self.assertEqual(first, again)
|
||||
self.assertNotEqual(first, other)
|
||||
self.assertNotIn(key_a, commitment)
|
||||
self.assertEqual(len(commitment), 64)
|
||||
|
||||
def test_source_hash_is_stable_hex(self) -> None:
|
||||
first = environment_source_sha256()
|
||||
self.assertEqual(first, environment_source_sha256())
|
||||
self.assertEqual(len(first), 64)
|
||||
self.assertTrue(set(first) <= set(string.hexdigits.lower()))
|
||||
|
||||
def test_task_data_carries_replay_identity_without_secret_key(self) -> None:
|
||||
taskset = RedactionTaskset(RedactionConfig(num_tasks=2, split="dev"))
|
||||
tasks = taskset.load()
|
||||
self.assertEqual([task.data.name for task in tasks], ["dev-00000", "dev-00001"])
|
||||
self.assertNotEqual(tasks[0].data.seed, tasks[1].data.seed)
|
||||
self.assertEqual(tasks[0].data.source_sha256, environment_source_sha256())
|
||||
self.assertEqual(len(tasks[0].data.seed_commitment), 64)
|
||||
self.assertEqual(tasks[0].data.environment_version, "0.2.0")
|
||||
|
||||
def test_scoring_records_provenance_and_oracle_metrics(self) -> None:
|
||||
task = RedactionTaskset(RedactionConfig(num_tasks=1, split="dev")).load()[0]
|
||||
trace = SimpleNamespace(last_reply=reply_for(ORACLE_RULES), info={})
|
||||
metrics = asyncio.run(task.scan(trace))
|
||||
self.assertEqual(metrics["recall"], 1.0)
|
||||
self.assertEqual(metrics["precision"], 1.0)
|
||||
self.assertEqual(metrics["clean"], 1.0)
|
||||
self.assertEqual(metrics["validity"], 1.0)
|
||||
self.assertEqual(trace.info["environment"]["source_sha256"], environment_source_sha256())
|
||||
self.assertEqual(trace.info["redaction_errors"], [])
|
||||
|
||||
def test_invalidity_multiplies_all_rewards_to_zero(self) -> None:
|
||||
task = RedactionTaskset(RedactionConfig(num_tasks=1, split="dev")).load()[0]
|
||||
trace = SimpleNamespace(
|
||||
metrics={"recall": 1.0, "precision": 1.0, "clean": 1.0, "validity": 0.0}
|
||||
)
|
||||
self.assertEqual(asyncio.run(task.recall(trace)), 0.0)
|
||||
self.assertEqual(asyncio.run(task.precision(trace)), 0.0)
|
||||
self.assertEqual(asyncio.run(task.gate(trace)), 0.0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user