121 lines
5.0 KiB
Python
121 lines
5.0 KiB
Python
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()
|