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()