Files
nucleic-purpose-classifier/tests/test_label_swe_chat_prose.py
T

169 lines
7.2 KiB
Python

import json
import sys
import unittest
from pathlib import Path
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import label_nucleic_prompts as base
import label_swe_chat_prose as prose_label
from purpose_data import prompt_hash
PROSE = "Rewired the settings screen to the new theme tokens."
TEACHER_PROMPT = "Make the settings screen use the new theme."
class LabelSWEChatProseTests(unittest.TestCase):
def value(self, **overrides):
value = {
"schemaVersion": 1,
"repoID": "repo",
"userID": "user",
"sessionID": "session",
"sourceTurnIDs": ["one", "two"],
"prose": PROSE,
"teacherPrompt": TEACHER_PROMPT,
"endsInQuestion": False,
"tailBiased": False,
}
value["proseHash"] = prompt_hash(value["prose"])
value["teacherPromptHash"] = prompt_hash(value["teacherPrompt"])
value.update(overrides)
return value
def candidate(self, **overrides):
value = self.value(**overrides)
line = base.SourceLine(1, json.dumps(value), "line-hash")
return prose_label.candidates([line])[0]
def decision(self, candidate, recoverable, **overrides):
item = {
"id": candidate.id, "keep": True, "junkReason": None,
"purpose": "frontendImpl", "secondary": None, "mixed": False,
"difficulty": 0.4, "slice": "boundary", "lang": "en",
"recoverableFromProse": recoverable,
}
item.update(overrides)
return {"items": [item]}
def test_candidate_rejects_a_prose_hash_mismatch(self):
with self.assertRaisesRegex(ValueError, "proseHash"):
prose_label.candidates(
[base.SourceLine(1, json.dumps(self.value(proseHash="bad")), "line-hash")]
)
def test_candidate_rejects_a_teacher_prompt_hash_mismatch(self):
with self.assertRaisesRegex(ValueError, "teacherPromptHash"):
prose_label.candidates(
[base.SourceLine(1, json.dumps(self.value(teacherPromptHash="bad")), "line-hash")]
)
def test_candidate_requires_the_runtime_signals(self):
for field in ("endsInQuestion", "tailBiased"):
with self.subTest(field), self.assertRaisesRegex(ValueError, field):
prose_label.candidates(
[base.SourceLine(1, json.dumps(self.value(**{field: None})), "line-hash")]
)
def test_candidate_rejects_duplicate_normalized_prose(self):
lines = [
base.SourceLine(1, json.dumps(self.value()), "a"),
base.SourceLine(2, json.dumps(self.value(prose=f" {PROSE.upper()} ",
proseHash=prompt_hash(PROSE))), "b"),
]
with self.assertRaisesRegex(ValueError, "duplicate normalized prose"):
prose_label.candidates(lines)
def test_student_text_is_the_prose_never_the_user_prompt(self):
candidate = self.candidate()
decision = self.decision(candidate, True)["items"][0]
record = prose_label.record_from_decision(candidate, decision)
self.assertEqual(PROSE, record["prompt"])
self.assertNotIn(TEACHER_PROMPT, json.dumps(record))
self.assertEqual({"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"},
set(record))
def test_prose_that_needs_the_user_prompt_becomes_vague_eval(self):
candidate = self.candidate()
decisions = prose_label.validate_decisions(
[candidate], self.decision(candidate, False)
)
self.assertEqual("vague-eval", decisions[0][1]["slice"])
def test_unrecoverable_mixed_decision_becomes_single_purpose_vague_eval(self):
candidate = self.candidate()
decision = self.decision(
candidate, False, secondary="backendImpl", mixed=True, slice="mixed"
)["items"][0]
record = prose_label.record_from_decision(candidate, decision)
self.assertEqual("vague-eval", record["slice"])
self.assertFalse(record["mixed"])
self.assertIsNone(record["secondary"])
def test_decision_requires_the_recoverability_verdict(self):
candidate = self.candidate()
payload = self.decision(candidate, None)
with self.assertRaisesRegex(ValueError, "recoverableFromProse"):
prose_label.validate_decisions([candidate], payload)
def test_rejected_decision_must_not_claim_recoverability(self):
candidate = self.candidate()
payload = self.decision(
candidate, True, keep=False, junkReason="only offers work", purpose=None,
secondary=None, mixed=None, difficulty=None, slice=None, lang=None,
)
with self.assertRaisesRegex(ValueError, "recoverableFromProse=null"):
prose_label.validate_decisions([candidate], payload)
def test_state_and_audit_carry_no_user_prompt_text(self):
candidate = self.candidate()
state = prose_label._state(
candidate, status="labeled", record=None, reason=None, recoverable=True
)
self.assertNotIn(TEACHER_PROMPT, json.dumps(state))
self.assertEqual(prompt_hash(TEACHER_PROMPT), state["teacherPromptHash"])
self.assertTrue(set(prose_label.AUDIT_FIELDS) <= set(state))
self.assertNotIn("record", prose_label.AUDIT_FIELDS)
def test_labeling_prompt_quotes_both_texts_and_names_the_labeled_one(self):
candidate = self.candidate(endsInQuestion=True, tailBiased=True)
text = prose_label.labeling_prompt([candidate], 8_000, 24_000)
payload = json.loads(text.split("<input_json>\n", 1)[1].split("\n</input_json>", 1)[0])
item = payload["items"][0]
self.assertEqual(PROSE, item["agent_reply_prose"])
self.assertEqual(TEACHER_PROMPT, item["preceding_user_message"])
self.assertTrue(item["reply_was_truncated_to_its_tail"])
self.assertIn("Label only agent_reply_prose", text)
# The must-not-fire rule from CONTEXT_SWITCH.md §4 lives in the teacher prompt.
self.assertIn("only *asks* about or *offers* work", text)
def test_response_schema_requires_the_extra_verdict_field(self):
schema = prose_label.response_schema([self.candidate()])
item = schema["properties"]["items"]["items"]
self.assertIn("recoverableFromProse", item["properties"])
self.assertIn("recoverableFromProse", item["required"])
def test_batches_respect_size_and_character_budgets(self):
candidates = [
self.candidate(prose=f"Reply number {index}.",
proseHash=prompt_hash(f"Reply number {index}."))
for index in range(5)
]
batched = prose_label.batches(
candidates, batch_size=2, batch_chars=100_000,
max_prose_chars=8_000, max_prompt_chars=24_000,
)
self.assertEqual([2, 2, 1], [len(batch) for batch in batched])
with self.assertRaisesRegex(ValueError, "above --batch-chars"):
prose_label.batches(
candidates, batch_size=2, batch_chars=10,
max_prose_chars=8_000, max_prompt_chars=24_000,
)
if __name__ == "__main__":
unittest.main()