| from __future__ import annotations |
|
|
| from pathlib import Path |
| import tempfile |
| import unittest |
|
|
| from agent_harness.protocol_experiment import ( |
| ProtocolToolHarness, |
| ProtocolWorkspace, |
| _identity, |
| protocol_tool_definitions, |
| ) |
| from agent_harness.repair_experiment import PatchOutputError |
| from agent_harness.specs import ( |
| load_edit_interfaces, |
| load_experiments, |
| load_harnesses, |
| load_models, |
| load_task_split, |
| load_tasks, |
| ) |
| from agent_harness.study2_experiment import tokenizer_for |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| class ProtocolExperimentTests(unittest.TestCase): |
| def setUp(self) -> None: |
| self.task = load_tasks(ROOT)["TASK_CR_001"] |
| self.interfaces = load_edit_interfaces(ROOT) |
|
|
| def test_frozen_matrix_has_540_cells(self) -> None: |
| experiment = load_experiments(ROOT)["E09"] |
| split = load_task_split(ROOT / "tasks" / "splits" / "study3_protocol.txt") |
| self.assertEqual(len(split), 60) |
| self.assertEqual(experiment.cells_per_task(), 9) |
| self.assertEqual(experiment.cells_per_task() * len(split), 540) |
| self.assertEqual(experiment.edit_interface_ids, ("P001", "P002", "P003")) |
| self.assertEqual(experiment.model_ids, ("M002", "M003", "M004")) |
|
|
| def test_identity_records_ancillary_seed_and_context(self) -> None: |
| identity = _identity( |
| load_experiments(ROOT)["E12"], |
| load_tasks(ROOT)["TASK_S4_R001_003"], |
| self.interfaces["P002"], |
| load_models(ROOT)["M002"], |
| "a" * 40, |
| seed=2, |
| context_budget=16384, |
| retrieval_harness=load_harnesses(ROOT)["H007"], |
| ) |
| self.assertEqual(identity.seed, 2) |
| self.assertEqual(identity.context_budget, 16384) |
|
|
| def test_every_protocol_model_has_a_frozen_tokenizer(self) -> None: |
| models = load_models(ROOT) |
| for model_id in ("M002", "M003", "M004"): |
| tokenizer = tokenizer_for(models[model_id]) |
| self.assertTrue(tokenizer.path.is_file()) |
| self.assertGreater(tokenizer.count("protocol compatibility"), 0) |
|
|
| def test_structured_edits_create_canonical_final_diffs(self) -> None: |
| with tempfile.TemporaryDirectory() as temporary: |
| tree = Path(temporary) |
| target = tree / "example.go" |
| target.write_text("package example\n\nconst value = 1\n", encoding="utf-8") |
| workspace = ProtocolWorkspace(tree, ("example.go",), self.task, 2) |
| replaced = workspace.replace_text("example.go", "value = 1", "value = 2") |
| self.assertTrue(replaced["accepted"]) |
| self.assertIn("-const value = 1", workspace.final_patch()) |
| self.assertIn("+const value = 2", workspace.final_patch()) |
|
|
| written = workspace.write_file( |
| "example.go", "package example\n\nconst value = 3\n" |
| ) |
| self.assertTrue(written["accepted"]) |
| self.assertIn("+const value = 3", workspace.final_patch()) |
|
|
| def test_exact_replace_rejects_ambiguous_source(self) -> None: |
| with tempfile.TemporaryDirectory() as temporary: |
| tree = Path(temporary) |
| (tree / "example.go").write_text("x := 1\nx := 1\n", encoding="utf-8") |
| workspace = ProtocolWorkspace(tree, ("example.go",), self.task, 2) |
| with self.assertRaisesRegex(ValueError, "exactly once"): |
| workspace.replace_text("example.go", "x := 1", "x := 2") |
|
|
| def test_each_arm_exposes_only_its_assigned_edit_tool(self) -> None: |
| for interface in self.interfaces.values(): |
| names = [ |
| item["function"]["name"] |
| for item in protocol_tool_definitions(interface, self.task) |
| ] |
| self.assertEqual( |
| set(names), {"read_file", interface.edit_tool, "run_tests", "finish"} |
| ) |
|
|
| def test_retrieval_protocol_preserves_one_tool_signature(self) -> None: |
| harnesses = load_harnesses(ROOT) |
| for harness_id in ("H000", "H007"): |
| names = [ |
| item["function"]["name"] |
| for item in protocol_tool_definitions( |
| self.interfaces["P003"], self.task, harnesses[harness_id] |
| ) |
| ] |
| self.assertEqual( |
| names, ["search_code", "read_file", "write_file", "run_tests", "finish"] |
| ) |
| oracle_names = [ |
| item["function"]["name"] |
| for item in protocol_tool_definitions( |
| self.interfaces["P002"], self.task, harnesses["H018"] |
| ) |
| ] |
| self.assertNotIn("search_code", oracle_names) |
|
|
| def test_specialized_protocol_exposes_separate_search_actions(self) -> None: |
| names = [ |
| item["function"]["name"] |
| for item in protocol_tool_definitions( |
| self.interfaces["P002"], self.task, load_harnesses(ROOT)["H011"] |
| ) |
| ] |
| self.assertEqual( |
| names[:5], |
| ["search_exact", "search_lexical", "search_syntax", "search_dense", "search_graph"], |
| ) |
|
|
| def test_calling_an_unassigned_edit_tool_is_rejected(self) -> None: |
| with tempfile.TemporaryDirectory() as temporary: |
| tree = Path(temporary) |
| (tree / "example.go").write_text("package example\n", encoding="utf-8") |
| workspace = ProtocolWorkspace(tree, ("example.go",), self.task, 2) |
| tools = ProtocolToolHarness(self.interfaces["P002"], workspace) |
| with self.assertRaisesRegex(ValueError, "unavailable"): |
| tools.execute("apply_patch", {"patch": "not used"}) |
|
|
| def test_malformed_raw_diff_uses_the_scored_tool_error_domain(self) -> None: |
| with tempfile.TemporaryDirectory() as temporary: |
| tree = Path(temporary) |
| (tree / "example.go").write_text("package example\n", encoding="utf-8") |
| workspace = ProtocolWorkspace(tree, ("example.go",), self.task, 2) |
| tools = ProtocolToolHarness(self.interfaces["P001"], workspace) |
| with self.assertRaisesRegex(PatchOutputError, "no modified repository path"): |
| tools.execute("apply_patch", {"patch": "not a unified diff"}) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|