Files

314 lines
13 KiB
Python
Raw Permalink Normal View History

import json
import hashlib
import sys
import tempfile
import unittest
from pathlib import Path
from unittest import mock
MODULE_DIR = Path(__file__).resolve().parents[1]
if str(MODULE_DIR) not in sys.path:
sys.path.insert(0, str(MODULE_DIR))
import purpose_data
import rebuild_sol_high
def source_record(prompt: str, purpose: str = "backendImpl") -> dict:
return {
"prompt": prompt,
"purpose": purpose,
"secondary": None,
"mixed": False,
"difficulty": 0.4,
"slice": "core",
"lang": "en",
}
def raw_hash(raw: str) -> str:
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
class RebuildSolHighTests(unittest.TestCase):
def test_snapshot_locks_all_inputs_and_status_starts_incomplete(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
data = root / "data"
data.mkdir()
first = data / "first.jsonl"
second = data / "second.jsonl"
purpose_data.write_jsonl(first, [source_record("Build the API")])
purpose_data.write_jsonl(second, [source_record("Fix the layout", "frontendImpl")])
fixtures = root / "fixtures.json"
purpose_data.write_json(
fixtures,
[{"prompt": "Explain this module", "purpose": "review"}],
)
history = root / "history.jsonl"
purpose_data.write_jsonl(history, [{"prompt": "Add a cache"}])
swe = root / "swe.jsonl"
purpose_data.write_jsonl(swe, [{"prompt": "Debug the crash"}])
stage = root / "stage"
with mock.patch.object(
rebuild_sol_high, "PUBLIC_SOURCES", (first, second)
), mock.patch.object(rebuild_sol_high, "FIXTURES", fixtures):
config = rebuild_sol_high.snapshot(
stage=stage,
history_input=history,
swe_input=swe,
overwrite_stage=False,
)
result = rebuild_sol_high.status(stage)
self.assertEqual(
config["teacher"],
{"model": "gpt-5.6-sol", "reasoningEffort": "high"},
)
self.assertEqual(result["public"]["input"], 2)
self.assertEqual(result["fixtures"]["input"], 1)
self.assertFalse(result["complete"])
with (stage / "history.unlabeled.jsonl").open("a", encoding="utf-8") as handle:
handle.write('{"prompt":"changed"}\n')
with self.assertRaisesRegex(purpose_data.DataError, "input changed"):
rebuild_sol_high.load_config(stage)
def test_status_requires_one_terminal_decision_per_input_line(self):
with tempfile.TemporaryDirectory() as directory:
stage = Path(directory)
inputs = {}
for job in rebuild_sol_high.jobs(stage):
purpose_data.write_jsonl(job.input, [{"prompt": f"{job.name} prompt"}])
inputs[job.input.name] = {
"records": 1,
"sha256": purpose_data.file_sha256(job.input),
}
state = {
"schemaVersion": 1,
"sourceLine": 1,
"sourceLineHash": "unused-by-status",
"promptHash": "unused-by-status",
"sessionID": None,
"status": "labeled",
"reason": None,
"record": source_record(f"{job.name} prompt"),
}
purpose_data.write_jsonl(job.state, [state])
purpose_data.write_jsonl(stage / "public-map.jsonl", [])
purpose_data.write_json(
stage / "workflow.json",
{
"schemaVersion": 1,
"teacher": {
"model": "gpt-5.6-sol",
"reasoningEffort": "high",
},
"inputs": inputs,
"sourceInputs": {},
},
)
result = rebuild_sol_high.status(stage)
self.assertTrue(result["complete"])
self.assertTrue(
all(
result[name]["labeled"] == 1
for name in ("public", "fixtures", "history", "swe")
)
)
def test_training_commands_are_from_base_for_both_tiers(self):
commands = rebuild_sol_high.train_commands("$PY")
self.assertIn("train.py", commands)
self.assertIn("train_deep_mlx.py", commands)
self.assertIn("--variant base", commands)
self.assertEqual(commands.count("--dataset-dir"), 2)
self.assertNotIn("--resume-from", commands)
self.assertNotIn("--model ", commands)
def test_training_commands_default_to_invoking_interpreter(self):
args = rebuild_sol_high.build_parser().parse_args(["train-commands"])
self.assertEqual(args.python, sys.executable)
self.assertTrue(
rebuild_sol_high.train_commands(args.python).startswith(
f'"{sys.executable}" ml/purpose-classifier/train.py'
)
)
def test_complete_stage_promotes_public_and_combined_datasets(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
data = root / "ml" / "purpose-classifier" / "data"
data.mkdir(parents=True)
script_dir = data.parent
first = data / "purpose-prompts.jsonl"
second = data / "purpose-prompts-round2.jsonl"
records = [
source_record(
f"Implement sample endpoint number {index} with stable pagination",
purpose_data.LABELS[index % len(purpose_data.LABELS)],
)
for index in range(80)
]
purpose_data.write_jsonl(first, records[:40])
purpose_data.write_jsonl(second, records[40:])
purpose_data.write_json(
data / "generation-manifest.json",
{
"schemaVersion": 1,
"canonicalFiles": [first.name, second.name],
"derivedBatchFiles": [],
"generations": [],
"limitations": [],
},
)
purpose_data.write_json(
data / "curation-review-v1.json",
{
"schemaVersion": 1,
"datasetVersion": "purpose-dataset-v1",
"semanticDuplicateReview": {"status": "complete"},
"humanLabelAndDifficultyReview": {"status": "complete"},
},
)
fixtures = root / "Tests" / "Fixtures" / "purpose-prompts.json"
purpose_data.write_json(
fixtures,
[{"prompt": "Plan the cache migration", "purpose": "planning"}],
)
history = root / "history.jsonl"
purpose_data.write_jsonl(history, [{"prompt": "Add a private cache layer"}])
swe = root / "swe.jsonl"
duplicate_words = [f"duplicate{index}" for index in range(100)]
first_duplicate = " ".join(duplicate_words)
duplicate_words[50] = "replacement"
second_duplicate = " ".join(duplicate_words)
purpose_data.write_jsonl(
swe,
[
{"prompt": "Diagnose a unique worker crash"},
{"prompt": first_duplicate},
{"prompt": second_duplicate},
],
)
stage = script_dir / ".artifacts" / "sol-high-reset"
public_dataset = script_dir / ".artifacts" / "dataset-public"
combined_dataset = script_dir / ".artifacts" / "dataset-v1"
patches = (
mock.patch.object(rebuild_sol_high, "REPOSITORY_ROOT", root),
mock.patch.object(rebuild_sol_high, "DATA_DIR", data),
mock.patch.object(rebuild_sol_high, "FIXTURES", fixtures),
mock.patch.object(rebuild_sol_high, "PUBLIC_SOURCES", (first, second)),
mock.patch.object(
rebuild_sol_high,
"PUBLIC_DATASET_DESTINATION",
public_dataset,
),
mock.patch.object(
rebuild_sol_high,
"COMBINED_DATASET_DESTINATION",
combined_dataset,
),
)
with patches[0], patches[1], patches[2], patches[3], patches[4], patches[5]:
rebuild_sol_high.snapshot(
stage=stage,
history_input=history,
swe_input=swe,
overwrite_stage=False,
)
for job in rebuild_sol_high.jobs(stage):
input_lines = job.input.read_text(encoding="utf-8").splitlines()
states = []
for line, raw in enumerate(input_lines, 1):
prompt = json.loads(raw)["prompt"]
original = next(
(item for item in records if item["prompt"] == prompt),
None,
)
record = original or source_record(
prompt,
"writing" if prompt == second_duplicate else "backendImpl",
)
states.append(
{
"schemaVersion": 1,
"sourceLine": line,
"sourceLineHash": raw_hash(raw),
"promptHash": "test",
"sessionID": None,
"status": "labeled",
"reason": None,
"record": record,
}
)
purpose_data.write_jsonl(job.state, states)
result = rebuild_sol_high.promote(
stage,
rebuild_sol_high.CONFIRMATION,
excluded_real_lines=[3],
)
self.assertEqual(80, result["publicLabeled"])
self.assertEqual(3, result["realLabeled"])
self.assertEqual(3, result["realReviewExclusions"][0]["sourceLine"])
self.assertTrue((combined_dataset / "train.jsonl").is_file())
self.assertGreater(
len(purpose_data.load_jsonl(combined_dataset / "train.jsonl")),
len(purpose_data.load_jsonl(public_dataset / "train.jsonl")),
)
promoted_manifest = json.loads(
(data / "dataset-v1-manifest.json").read_text(encoding="utf-8")
)
self.assertEqual(
"purpose-dataset-sol-high-v2",
promoted_manifest["datasetVersion"],
)
combined_manifest = json.loads(
(combined_dataset / "manifest.json").read_text(encoding="utf-8")
)
self.assertEqual(
result["realReviewExclusions"],
combined_manifest["reviewedRealExclusions"],
)
def test_promotion_requires_exact_confirmation(self):
with self.assertRaisesRegex(purpose_data.DataError, "promotion requires"):
rebuild_sol_high.promote(Path("unused"), None)
def test_reviewed_real_line_exclusions_are_auditable(self):
records = [
source_record("Plan the cache migration", "planning"),
source_record("Review the cache migration", "review"),
]
retained, review = rebuild_sol_high._exclude_reviewed_real_lines(records, [1])
self.assertEqual([records[1]], retained)
self.assertEqual(1, review[0]["sourceLine"])
self.assertEqual("planning", review[0]["purpose"])
self.assertEqual(
"reviewed-near-duplicate-label-conflict", review[0]["reason"]
)
self.assertEqual(64, len(review[0]["promptHash"]))
def test_reviewed_real_line_exclusions_reject_invalid_lines(self):
records = [source_record("Plan the cache migration", "planning")]
with self.assertRaisesRegex(purpose_data.DataError, "duplicate line"):
rebuild_sol_high._exclude_reviewed_real_lines(records, [1, 1])
with self.assertRaisesRegex(purpose_data.DataError, "outside"):
rebuild_sol_high._exclude_reviewed_real_lines(records, [2])
if __name__ == "__main__":
unittest.main()