agent-harness / tests /test_robustness_experiment.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
2.48 kB
from __future__ import annotations
import unittest
from agent_harness.robustness_experiment import (
RobustnessExperimentError,
distractor_severity,
robustness_shift,
synthetic_distractor_sources,
)
from agent_harness.specs import TaskSpec
from agent_harness.syntax_index import parse_go_file
def task() -> TaskSpec:
return TaskSpec(
schema_version=1,
task_id="TASK_CR_999",
repository_url="https://gitlab.com/gitlab-org/gitlab-runner.git",
base_commit="0" * 40,
gold_commit="1" * 40,
language="go",
statement="Cache the role ARN for the docker autoscaler executor.",
gold_patch="gold.patch",
test_patch="test.patch",
gold_files=("executors/docker/autoscaler.go",),
gold_symbols=("executors/docker/autoscaler.go::RoleARN",),
fail_to_pass_tests=("TestRoleARN",),
pass_to_pass_tests=("TestExisting",),
difficulty="test",
provenance="unit test",
validation_status="end_to_end_ready",
)
class RobustnessExperimentTests(unittest.TestCase):
def test_seed_to_dose_mapping_is_frozen(self) -> None:
self.assertEqual([distractor_severity(seed) for seed in (0, 1, 2)], [1, 5, 10])
with self.assertRaises(RobustnessExperimentError):
distractor_severity(3)
def test_distractors_are_nested_deterministic_and_parseable(self) -> None:
one = synthetic_distractor_sources(task(), 1)
ten = synthetic_distractor_sources(task(), 10)
self.assertEqual(one, ten[:1])
self.assertEqual(len({source.path for source in ten}), 10)
self.assertTrue(all(parse_go_file(source.path, source.text) for source in ten))
self.assertTrue(all("role ARN" in source.text for source in ten))
def test_shift_uses_censoring_for_missing_gold(self) -> None:
baseline = {"file_recall_at_10": 1.0, "mrr": 0.5, "ndcg_at_10": 0.7, "first_gold_rank": 2}
perturbed = {"file_recall_at_10": 0.0, "mrr": 0.0, "ndcg_at_10": 0.0, "first_gold_rank": None}
shift = robustness_shift(
baseline,
perturbed,
["other.go", "gold.go"],
["distractor.go"],
["gold.go"],
missing_rank=201,
)
self.assertEqual(shift["first_gold_rank_displacement_censored"], 199)
self.assertTrue(shift["lost_all_top_10_gold"])
self.assertEqual(shift["baseline_gold_top_10_retention"], 0.0)