Harden redaction-pressure reward and provenance
arena-environments / validate (3.11) (push) Successful in 1m8s
arena-environments / validate (3.12) (push) Successful in 31s

This commit is contained in:
2026-08-19 01:26:21 -07:00
parent 5a99eb86a5
commit d691acf131
16 changed files with 897 additions and 163 deletions
@@ -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()