agent-harness / tests /test_retrieval.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
2.07 kB
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()