diff --git a/label_swe_chat_prompts.py b/label_swe_chat_prompts.py index 8de36a5..4934d70 100644 --- a/label_swe_chat_prompts.py +++ b/label_swe_chat_prompts.py @@ -306,6 +306,16 @@ def _append(path: Path, states: Sequence[dict[str, Any]]) -> None: os.fsync(handle.fileno()) +def record_from_decision(candidate: Candidate, decision: dict[str, Any]) -> dict[str, Any]: + """Construct the exact seven-field student record from a teacher decision.""" + record = {field: decision[field] for field in SOURCE_FIELDS - {"prompt"}} + record["prompt"] = candidate.prompt + if not decision["recoverableFromFirst"]: + record["slice"] = "vague-eval" + validate_source_record(record, candidate.id) + return record + + def run(args: argparse.Namespace) -> dict[str, int]: lines = base.source_lines(args.input) source = candidates(lines) @@ -335,10 +345,7 @@ def run(args: argparse.Namespace) -> dict[str, int]: newly: list[dict[str, Any]] = [] for candidate, decision in decisions: if decision["keep"]: - record = {field: decision[field] for field in SOURCE_FIELDS} - if not decision["recoverableFromFirst"]: - record["slice"] = "vague-eval" - validate_source_record(record, candidate.id) + record = record_from_decision(candidate, decision) newly.append(_state(candidate, status="labeled", record=record, reason=None, recoverable=decision["recoverableFromFirst"])) else: newly.append(_state(candidate, status="rejected", record=None, reason=f"semantic_junk:{decision['junkReason'].strip()}", recoverable=None)) diff --git a/tests/test_label_swe_chat_prompts.py b/tests/test_label_swe_chat_prompts.py index ca2d14c..8639dbe 100644 --- a/tests/test_label_swe_chat_prompts.py +++ b/tests/test_label_swe_chat_prompts.py @@ -83,6 +83,13 @@ class LabelSWEChatPromptsTests(unittest.TestCase): with self.assertRaisesRegex(ValueError, "rather than truncating"): label_swe_chat_prompts.response_for_labeling("x" * 1_001, 1_000) + def test_record_uses_only_the_first_user_message_as_student_text(self): + candidate = self.candidate() + decision = self.decision(candidate, True)["items"][0] + record = label_swe_chat_prompts.record_from_decision(candidate, decision) + self.assertEqual(candidate.prompt, record["prompt"]) + self.assertNotIn(candidate.response, record.values()) + if __name__ == "__main__": unittest.main()