71 lines
2.5 KiB
Python
71 lines
2.5 KiB
Python
import json
|
|||
|
|
import sys
|
||
|
|
import tempfile
|
||
|
|
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_prompts
|
||
|
|
from purpose_data import prompt_hash
|
||
|
|
|
||
|
|
|
||
|
|
class LabelSWEChatPromptsTests(unittest.TestCase):
|
||
|
|
def candidate(self):
|
||
|
|
value = {
|
||
|
|
"schemaVersion": 1,
|
||
|
|
"repoID": "repo",
|
||
|
|
"userID": "user",
|
||
|
|
"sessionID": "session",
|
||
|
|
"sourceTurnIDs": ["one", "two", "three"],
|
||
|
|
"prompt": "What is making this test fail?",
|
||
|
|
"teacherContext": ["It fails only on CI.", "Please diagnose it."],
|
||
|
|
}
|
||
|
|
value["promptHash"] = prompt_hash(value["prompt"])
|
||
|
|
value["contextPromptHashes"] = [prompt_hash(item) for item in value["teacherContext"]]
|
||
|
|
line = base.SourceLine(1, json.dumps(value), "line-hash")
|
||
|
|
return label_swe_chat_prompts.candidates([line])[0]
|
||
|
|
|
||
|
|
def decision(self, candidate, recoverable):
|
||
|
|
return {
|
||
|
|
"items": [{
|
||
|
|
"id": candidate.id, "keep": True, "junkReason": None,
|
||
|
|
"purpose": "debugging", "secondary": None, "mixed": False,
|
||
|
|
"difficulty": 0.6, "slice": "boundary", "lang": "en",
|
||
|
|
"recoverableFromFirst": recoverable,
|
||
|
|
}]
|
||
|
|
}
|
||
|
|
|
||
|
|
def test_context_hashes_are_checked_and_context_dependent_labels_become_vague_eval(self):
|
||
|
|
candidate = self.candidate()
|
||
|
|
decisions = label_swe_chat_prompts.validate_decisions(
|
||
|
|
[candidate], self.decision(candidate, False)
|
||
|
|
)
|
||
|
|
self.assertEqual("vague-eval", decisions[0][1]["slice"])
|
||
|
|
self.assertNotIn("It fails only on CI.", candidate.prompt)
|
||
|
|
|
||
|
|
def test_candidate_rejects_context_hash_mismatch(self):
|
||
|
|
candidate = self.candidate()
|
||
|
|
value = json.loads(candidate.line.raw)
|
||
|
|
value["contextPromptHashes"][0] = "bad"
|
||
|
|
with self.assertRaisesRegex(ValueError, "contextPromptHashes"):
|
||
|
|
label_swe_chat_prompts.candidates(
|
||
|
|
[base.SourceLine(1, json.dumps(value), "line-hash")]
|
||
|
|
)
|
||
|
|
|
||
|
|
def test_state_never_carries_later_text(self):
|
||
|
|
candidate = self.candidate()
|
||
|
|
state = label_swe_chat_prompts._state(
|
||
|
|
candidate, status="labeled", record=None, reason=None, recoverable=True
|
||
|
|
)
|
||
|
|
encoded = json.dumps(state)
|
||
|
|
self.assertNotIn("It fails only on CI.", encoded)
|
||
|
|
self.assertNotIn("Please diagnose it.", encoded)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|