244 lines
10 KiB
Python
244 lines
10 KiB
Python
import sys
|
|||
|
|
import unittest
|
||
|
|
from collections import Counter
|
||
|
|
from pathlib import Path
|
||
|
|
|
||
|
|
|
||
|
|
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||
|
|
sys.path.insert(0, str(MODULE_DIR))
|
||
|
|
|
||
|
|
import export_swe_chat_prose as prose_export
|
||
|
|
|
||
|
|
|
||
|
|
def row(session, turn, ordinal, content, **overrides):
|
||
|
|
value = {
|
||
|
|
"session_id": session,
|
||
|
|
"turn_id": turn,
|
||
|
|
"conversation_turn_number": ordinal,
|
||
|
|
"turn_number": ordinal,
|
||
|
|
"turn_type": "user_prompt" if ordinal % 2 == 0 else "assistant_response",
|
||
|
|
"role": "user" if ordinal % 2 == 0 else "assistant",
|
||
|
|
"is_conversational": True,
|
||
|
|
"is_continuation": False,
|
||
|
|
"content": content,
|
||
|
|
}
|
||
|
|
value.update(overrides)
|
||
|
|
return value
|
||
|
|
|
||
|
|
|
||
|
|
def select(rows, *, max_per_session=3):
|
||
|
|
"""Run both passes the way `export` does, over one in-memory row list."""
|
||
|
|
|
||
|
|
markers, funnel = prose_export.conversational_markers(rows)
|
||
|
|
pairs = prose_export.select_turn_ends(markers, funnel, max_per_session=max_per_session)
|
||
|
|
candidates = prose_export.attach_text(
|
||
|
|
rows, pairs, sessions={}, funnel=funnel,
|
||
|
|
character_limit=prose_export.DEFAULT_CHARACTER_LIMIT,
|
||
|
|
)
|
||
|
|
return candidates, funnel
|
||
|
|
|
||
|
|
|
||
|
|
class TurnEndSelectionTests(unittest.TestCase):
|
||
|
|
def test_takes_only_the_last_reply_of_a_multi_message_assistant_run(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Add the settings view."),
|
||
|
|
row("s1", "t1", 1, "Working on it."),
|
||
|
|
row("s1", "t2", 2, "Still working.", turn_type="assistant_response", role="assistant"),
|
||
|
|
row("s1", "t3", 3, "The settings view is done."),
|
||
|
|
]
|
||
|
|
candidates, funnel = select(rows)
|
||
|
|
self.assertEqual(1, len(candidates))
|
||
|
|
self.assertEqual("The settings view is done.", candidates[0].prose)
|
||
|
|
self.assertEqual("t3", candidates[0].pair.response_turn_id)
|
||
|
|
self.assertEqual(2, funnel["rejectedMidTurnAssistantResponse"])
|
||
|
|
|
||
|
|
def test_keeps_the_final_reply_when_the_session_ends(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Ship it."),
|
||
|
|
row("s1", "t1", 1, "Shipped. Anything else?"),
|
||
|
|
]
|
||
|
|
candidates, _ = select(rows)
|
||
|
|
self.assertEqual(["Shipped. Anything else?"], [c.prose for c in candidates])
|
||
|
|
|
||
|
|
def test_pairs_each_reply_with_the_prompt_that_opened_its_own_turn(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "First request."),
|
||
|
|
row("s1", "t1", 1, "First answer."),
|
||
|
|
row("s1", "t2", 2, "Second request."),
|
||
|
|
row("s1", "t3", 3, "Second answer."),
|
||
|
|
]
|
||
|
|
candidates, _ = select(rows)
|
||
|
|
self.assertEqual(
|
||
|
|
[("First request.", "First answer."), ("Second request.", "Second answer.")],
|
||
|
|
[(c.teacher_prompt, c.prose) for c in candidates],
|
||
|
|
)
|
||
|
|
|
||
|
|
def test_orders_by_conversation_ordinal_not_row_arrival(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t3", 3, "Second answer."),
|
||
|
|
row("s1", "t1", 1, "First answer."),
|
||
|
|
row("s1", "t2", 2, "Second request."),
|
||
|
|
row("s1", "t0", 0, "First request."),
|
||
|
|
]
|
||
|
|
candidates, _ = select(rows)
|
||
|
|
self.assertEqual(["First answer.", "Second answer."], [c.prose for c in candidates])
|
||
|
|
|
||
|
|
def test_drops_a_reply_with_no_preceding_user_prompt(self):
|
||
|
|
rows = [row("s1", "t1", 1, "Orphan reply.")]
|
||
|
|
candidates, funnel = select(rows)
|
||
|
|
self.assertEqual([], candidates)
|
||
|
|
self.assertEqual(1, funnel["rejectedNoPrecedingUserPrompt"])
|
||
|
|
|
||
|
|
def test_caps_replies_taken_from_one_session(self):
|
||
|
|
rows = []
|
||
|
|
for index in range(4):
|
||
|
|
rows.append(row("s1", f"u{index}", index * 2, f"Request {index}."))
|
||
|
|
rows.append(row("s1", f"a{index}", index * 2 + 1, f"Answer {index}."))
|
||
|
|
candidates, funnel = select(rows, max_per_session=2)
|
||
|
|
self.assertEqual(["Answer 0.", "Answer 1."], [c.prose for c in candidates])
|
||
|
|
self.assertEqual(2, funnel["rejectedSessionCap"])
|
||
|
|
|
||
|
|
def test_rejects_non_conversational_continuation_and_malformed_rows(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Request."),
|
||
|
|
row("s1", "t1", 1, "Tool output.", is_conversational=False),
|
||
|
|
row("s1", "t2", 2, "Request.", is_continuation=True),
|
||
|
|
row("s1", "t3", 3, "Reply.", conversation_turn_number=None),
|
||
|
|
row("s1", "t4", 4, "Reply.", turn_type="tool_call", role="assistant"),
|
||
|
|
]
|
||
|
|
_, funnel = select(rows)
|
||
|
|
self.assertEqual(1, funnel["rejectedNonConversational"])
|
||
|
|
self.assertEqual(1, funnel["rejectedContinuation"])
|
||
|
|
self.assertEqual(1, funnel["rejectedMalformedMarker"])
|
||
|
|
self.assertEqual(1, funnel["rejectedTurnTypeOrRole"])
|
||
|
|
|
||
|
|
def test_rejects_a_duplicate_turn_id_within_a_session(self):
|
||
|
|
rows = [row("s1", "t0", 0, "a"), row("s1", "t0", 2, "b")]
|
||
|
|
with self.assertRaises(prose_export.DataError):
|
||
|
|
prose_export.conversational_markers(rows)
|
||
|
|
|
||
|
|
|
||
|
|
class ProseExtractionTests(unittest.TestCase):
|
||
|
|
def test_student_text_is_the_runtime_extraction_not_the_raw_reply(self):
|
||
|
|
reply = "The migration is done.\n```swift\nstruct View {}\n```\nTests pass."
|
||
|
|
rows = [row("s1", "t0", 0, "Migrate it."), row("s1", "t1", 1, reply)]
|
||
|
|
candidates, _ = select(rows)
|
||
|
|
self.assertEqual("The migration is done.\nTests pass.", candidates[0].prose)
|
||
|
|
self.assertFalse(candidates[0].tail_biased)
|
||
|
|
|
||
|
|
def test_flags_prose_the_character_limit_actually_truncated(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Do it."),
|
||
|
|
row("s1", "t1", 1, "old words " * 6 + "Now the settings view."),
|
||
|
|
]
|
||
|
|
markers, funnel = prose_export.conversational_markers(rows)
|
||
|
|
pairs = prose_export.select_turn_ends(markers, funnel, max_per_session=3)
|
||
|
|
candidates = prose_export.attach_text(
|
||
|
|
rows, pairs, sessions={}, funnel=funnel, character_limit=25
|
||
|
|
)
|
||
|
|
self.assertEqual("Now the settings view.", candidates[0].prose)
|
||
|
|
self.assertTrue(candidates[0].tail_biased)
|
||
|
|
|
||
|
|
def test_drops_a_reply_that_is_entirely_code(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Show me the struct."),
|
||
|
|
row("s1", "t1", 1, "```swift\nstruct View {}\n```"),
|
||
|
|
]
|
||
|
|
candidates, funnel = select(rows)
|
||
|
|
self.assertEqual([], candidates)
|
||
|
|
self.assertEqual(1, funnel["rejectedNoProseAfterExtraction"])
|
||
|
|
|
||
|
|
def test_drops_empty_and_nul_bearing_content(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Request."),
|
||
|
|
row("s1", "t1", 1, " "),
|
||
|
|
row("s1", "t2", 2, "Request."),
|
||
|
|
row("s1", "t3", 3, "reply\x00"),
|
||
|
|
]
|
||
|
|
candidates, funnel = select(rows)
|
||
|
|
self.assertEqual([], candidates)
|
||
|
|
self.assertEqual(2, funnel["rejectedMalformedOrEmptyContent"])
|
||
|
|
self.assertEqual(2, funnel["rejectedMissingText"])
|
||
|
|
|
||
|
|
def test_records_the_question_signal_and_hashes_the_student_text(self):
|
||
|
|
rows = [
|
||
|
|
row("s1", "t0", 0, "Finish the backend."),
|
||
|
|
row("s1", "t1", 1, "Backend is done. Should I start the UI?"),
|
||
|
|
]
|
||
|
|
candidates, _ = select(rows)
|
||
|
|
record = candidates[0].json("abc123", prose_export.DEFAULT_CHARACTER_LIMIT)
|
||
|
|
self.assertTrue(record["endsInQuestion"])
|
||
|
|
self.assertEqual(["t0", "t1"], record["sourceTurnIDs"])
|
||
|
|
self.assertEqual("abc123", record["sourceRevision"])
|
||
|
|
self.assertEqual(
|
||
|
|
prose_export.prompt_hash("Backend is done. Should I start the UI?"),
|
||
|
|
record["proseHash"],
|
||
|
|
)
|
||
|
|
self.assertEqual("Backend is done. Should I start the UI?", record["prose"])
|
||
|
|
|
||
|
|
|
||
|
|
class CapAndDedupeTests(unittest.TestCase):
|
||
|
|
def make(self, prose, repo, user):
|
||
|
|
pair = prose_export.Pair("s", "t0", "t1", 1)
|
||
|
|
return prose_export.Candidate(
|
||
|
|
session_id="s", repo_id=repo, user_id=user, pair=pair,
|
||
|
|
teacher_prompt="Request.", prose=prose, tail_biased=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
def test_dedupes_repeated_sign_off_prose_across_sessions(self):
|
||
|
|
candidates, funnel = prose_export.cap_and_dedupe(
|
||
|
|
[
|
||
|
|
self.make("All tests pass.", "r1", "u1"),
|
||
|
|
self.make(" all TESTS pass. ", "r2", "u2"),
|
||
|
|
self.make("Renamed the module.", "r3", "u3"),
|
||
|
|
],
|
||
|
|
Counter(), max_per_repo=10, max_per_user=10,
|
||
|
|
)
|
||
|
|
self.assertEqual(["All tests pass.", "Renamed the module."], [c.prose for c in candidates])
|
||
|
|
self.assertEqual(1, funnel["rejectedDuplicateProse"])
|
||
|
|
|
||
|
|
def test_caps_repository_and_user_concentration(self):
|
||
|
|
candidates, funnel = prose_export.cap_and_dedupe(
|
||
|
|
[
|
||
|
|
self.make("One.", "r1", "u1"),
|
||
|
|
self.make("Two.", "r1", "u2"),
|
||
|
|
self.make("Three.", "r2", "u1"),
|
||
|
|
self.make("Four.", "r2", "u2"),
|
||
|
|
],
|
||
|
|
Counter(), max_per_repo=1, max_per_user=1,
|
||
|
|
)
|
||
|
|
self.assertEqual(["One.", "Four."], [c.prose for c in candidates])
|
||
|
|
self.assertEqual(1, funnel["rejectedRepoCap"])
|
||
|
|
self.assertEqual(1, funnel["rejectedUserCap"])
|
||
|
|
self.assertEqual(2, funnel["exportedCandidates"])
|
||
|
|
|
||
|
|
|
||
|
|
class ArgumentTests(unittest.TestCase):
|
||
|
|
def test_rejects_a_mutable_revision(self):
|
||
|
|
for revision in ["main", ""]:
|
||
|
|
with self.subTest(revision), self.assertRaises(prose_export.DataError):
|
||
|
|
prose_export.export(
|
||
|
|
conversations=[], sessions_path=[], revision=revision,
|
||
|
|
output=Path("/dev/null"), manifest_path=Path("/dev/null"),
|
||
|
|
max_per_repo=1, max_per_user=1, max_per_session=1, character_limit=1,
|
||
|
|
)
|
||
|
|
|
||
|
|
def test_rejects_non_positive_caps_and_limits(self):
|
||
|
|
for kwargs in [
|
||
|
|
{"max_per_repo": 0}, {"max_per_user": 0},
|
||
|
|
{"max_per_session": 0}, {"character_limit": 0},
|
||
|
|
]:
|
||
|
|
settings = {
|
||
|
|
"max_per_repo": 1, "max_per_user": 1,
|
||
|
|
"max_per_session": 1, "character_limit": 1, **kwargs,
|
||
|
|
}
|
||
|
|
with self.subTest(kwargs), self.assertRaises(prose_export.DataError):
|
||
|
|
prose_export.export(
|
||
|
|
conversations=[], sessions_path=[], revision="abc123",
|
||
|
|
output=Path("/dev/null"), manifest_path=Path("/dev/null"), **settings,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|