File size: 2,070 Bytes
d61821a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 | from __future__ import annotations
import unittest
from agent_harness.pilot import retrieval_metrics
from agent_harness.repository import SourceFile, chunk_file
from agent_harness.retrieval import BM25FuzzyRetriever, ExactRetriever, query_terms
class RetrievalTests(unittest.TestCase):
def setUp(self) -> None:
sources = [
SourceFile(
"cache/s3.go",
"package cache\nfunc newS3Client() { checksumValidation := whenRequired }\n",
),
SourceFile(
"commands/proxy.go",
"package commands\nfunc executeProxy() { waitForChild() }\n",
),
]
self.chunks = tuple(
chunk
for source in sources
for chunk in chunk_file(source, chunk_lines=120, overlap_lines=20, char_limit=16000)
)
def test_query_terms_remove_stop_words_and_split_identifiers(self) -> None:
terms = query_terms("Fix checksumValidation for the S3-compatible endpoint")
self.assertIn("checksum", terms)
self.assertIn("validation", terms)
self.assertIn("s3", terms)
self.assertNotIn("the", terms)
def test_exact_and_bm25_retrieve_relevant_file(self) -> None:
query = "S3 checksum validation fails for custom endpoints"
exact = ExactRetriever(self.chunks).retrieve(query, 10)
bm25 = BM25FuzzyRetriever(self.chunks).retrieve(query, 10)
self.assertEqual(exact[0].path, "cache/s3.go")
self.assertEqual(bm25[0].path, "cache/s3.go")
def test_retrieval_metrics(self) -> None:
metrics = retrieval_metrics(
["irrelevant.go", "gold_a.go", "gold_b.go"],
["gold_a.go", "gold_b.go"],
)
self.assertEqual(metrics["file_recall_at_1"], 0.0)
self.assertEqual(metrics["file_recall_at_5"], 1.0)
self.assertEqual(metrics["first_gold_rank"], 2)
self.assertEqual(metrics["mrr"], 0.5)
self.assertTrue(metrics["all_gold_in_top_10"])
if __name__ == "__main__":
unittest.main()
|