Skip to content

Commit 52db32f

Browse files
committed
test: add signal extraction tests for advisor verdict attachment
FakeDatasource now accepts advisor_verdicts dict, enabling end-to-end tests that verify verdicts are correctly attached to SignalCommit objects during extraction. 3 new tests: - verdict attached to matching (commit, signal_key) - verdict for wrong signal_key is not attached - verdict attached to test-track signals
1 parent bd297fa commit 52db32f

1 file changed

Lines changed: 105 additions & 4 deletions

File tree

aws/lambda/pytorch-auto-revert/pytorch_auto_revert/tests/test_signal_extraction.py

Lines changed: 105 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,9 +25,15 @@ def ts(base: datetime, minutes: int) -> datetime:
2525
class FakeDatasource(SignalExtractionDatasource):
2626
"""Test double for the datasource returning provided rows."""
2727

28-
def __init__(self, jobs: List[JobRow], tests: List[TestRow]):
28+
def __init__(
29+
self,
30+
jobs: List[JobRow],
31+
tests: List[TestRow],
32+
advisor_verdicts: Optional[dict] = None,
33+
):
2934
self._jobs = jobs
3035
self._tests = tests
36+
self._advisor_verdicts = advisor_verdicts or {}
3137

3238
def fetch_commits_in_time_range(
3339
self,
@@ -67,7 +73,7 @@ def fetch_tests_for_job_ids(
6773
return [r for r in self._tests if int(r.job_id) in ids]
6874

6975
def fetch_advisor_verdicts(self, **kwargs):
70-
return {}
76+
return self._advisor_verdicts
7177

7278

7379
def J(
@@ -124,9 +130,14 @@ class TestSignalExtraction(unittest.TestCase):
124130
def setUp(self) -> None:
125131
self.t0 = datetime(2025, 8, 20, 12, 0, 0)
126132

127-
def _extract(self, jobs: List[JobRow], tests: List[TestRow]):
133+
def _extract(
134+
self,
135+
jobs: List[JobRow],
136+
tests: List[TestRow],
137+
advisor_verdicts: Optional[dict] = None,
138+
):
128139
se = SignalExtractor(workflows=["trunk"], lookback_hours=24)
129-
se._datasource = FakeDatasource(jobs, tests)
140+
se._datasource = FakeDatasource(jobs, tests, advisor_verdicts)
130141
return se.extract()
131142

132143
def _find_job_signal(self, signals, wf: str, base: JobBaseName):
@@ -925,5 +936,95 @@ def test_no_job_test_signal_when_only_non_test_failures(self):
925936
self.assertIsNone(self._find_job_signal(signals, "trunk", f"{base} [test]"))
926937

927938

939+
def test_advisor_verdict_attached_to_correct_commit(self):
940+
"""Advisor verdict from datasource is attached to the matching (commit, signal)."""
941+
jobs = [
942+
J(
943+
sha="C2", run=200, job=1, attempt=1,
944+
started_at=ts(self.t0, 10), conclusion="failure",
945+
),
946+
J(
947+
sha="C1", run=100, job=2, attempt=1,
948+
started_at=ts(self.t0, 5), conclusion="success",
949+
),
950+
]
951+
base = jobs[0].base_name
952+
# Advisor says revert for C2 on the job signal
953+
advisor_verdicts = {
954+
("C2", base): ("revert", 0.95, self.t0),
955+
}
956+
signals = self._extract(jobs, tests=[], advisor_verdicts=advisor_verdicts)
957+
sig = self._find_job_signal(signals, "trunk", base)
958+
self.assertIsNotNone(sig)
959+
960+
# C2 should have advisor_result
961+
c2 = next(c for c in sig.commits if c.head_sha == "C2")
962+
self.assertIsNotNone(c2.advisor_result)
963+
self.assertEqual(c2.advisor_result.verdict.value, "revert")
964+
self.assertAlmostEqual(c2.advisor_result.confidence, 0.95)
965+
self.assertEqual(c2.advisor_result.signal_key, base)
966+
967+
# C1 should NOT have advisor_result
968+
c1 = next(c for c in sig.commits if c.head_sha == "C1")
969+
self.assertIsNone(c1.advisor_result)
970+
971+
def test_advisor_verdict_not_attached_to_wrong_signal(self):
972+
"""Advisor verdict for signal A is not attached to signal B."""
973+
jobs = [
974+
J(
975+
sha="C2", run=200, job=1, attempt=1,
976+
started_at=ts(self.t0, 10), conclusion="failure",
977+
),
978+
J(
979+
sha="C1", run=100, job=2, attempt=1,
980+
started_at=ts(self.t0, 5), conclusion="success",
981+
),
982+
]
983+
base = jobs[0].base_name
984+
# Advisor verdict is for a completely different signal key
985+
advisor_verdicts = {
986+
("C2", "some_other_signal"): ("not_related", 0.99, self.t0),
987+
}
988+
signals = self._extract(jobs, tests=[], advisor_verdicts=advisor_verdicts)
989+
sig = self._find_job_signal(signals, "trunk", base)
990+
self.assertIsNotNone(sig)
991+
992+
# C2 should NOT have advisor_result (wrong signal key)
993+
c2 = next(c for c in sig.commits if c.head_sha == "C2")
994+
self.assertIsNone(c2.advisor_result)
995+
996+
def test_advisor_verdict_on_test_signal(self):
997+
"""Advisor verdict is attached to test-track signals."""
998+
jobs = [
999+
J(
1000+
sha="C2", run=200, job=1, attempt=1,
1001+
started_at=ts(self.t0, 10), conclusion="failure",
1002+
rule="pytest failure",
1003+
),
1004+
J(
1005+
sha="C1", run=100, job=2, attempt=1,
1006+
started_at=ts(self.t0, 5), conclusion="success",
1007+
),
1008+
]
1009+
tests = [
1010+
T(job=1, run=200, attempt=1, file="test_foo.py",
1011+
name="test_bar", failure_runs=1, success_runs=0),
1012+
T(job=2, run=100, attempt=1, file="test_foo.py",
1013+
name="test_bar", failure_runs=0, success_runs=1),
1014+
]
1015+
test_key = "test_foo.py::test_bar"
1016+
advisor_verdicts = {
1017+
("C2", test_key): ("garbage", 0.92, self.t0),
1018+
}
1019+
signals = self._extract(jobs, tests, advisor_verdicts=advisor_verdicts)
1020+
sig = self._find_test_signal(signals, "trunk", test_key)
1021+
self.assertIsNotNone(sig)
1022+
1023+
c2 = next(c for c in sig.commits if c.head_sha == "C2")
1024+
self.assertIsNotNone(c2.advisor_result)
1025+
self.assertEqual(c2.advisor_result.verdict.value, "garbage")
1026+
self.assertEqual(c2.advisor_result.signal_key, test_key)
1027+
1028+
9281029
if __name__ == "__main__":
9291030
unittest.main()

0 commit comments

Comments
 (0)