themis / phase1 /eval /test_case_reference_resolution.py
vg15o2's picture
Moonley backend (HF Space build)
1d9bd9b
Raw
History Blame Contribute Delete
3.83 kB
import os
import re
import sys
import unittest
HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(HERE, "..", "scripts"))
from agent import resolve_case_reference
class FakeCorpus:
def __init__(self):
self.meta = {
"babu-specific": {
"case_name": "Babu Lal vs Hazari Lal Klshori Lal & Ors",
"neutral_citation": "1982 INSC 11",
"year": 1982,
},
"babu-other": {
"case_name": "Babu Lal vs State Of Uttar Pradesh",
"neutral_citation": "1964 INSC 40",
"year": 1964,
},
"bachan-death": {
"case_name": "Bachan Singh vs State Of Punjab",
"neutral_citation": "1980 INSC 120",
"year": 1980,
},
"bachan-service": {
"case_name": "Bachan Singh & Anr vs Union Of India & Ors",
"neutral_citation": "1972 INSC 85",
"year": 1972,
},
}
def is_retrieval_eligible(self, doc_id):
return str(doc_id) in self.meta
def _card(self, doc_id):
return {"doc_id": str(doc_id), **self.meta[str(doc_id)]}
@staticmethod
def _normal(value):
return re.sub(r"[^a-z0-9]+", " ", str(value).lower()).strip()
def identity_hits(self, query):
normalized = self._normal(query)
for doc_id, item in self.meta.items():
if normalized == self._normal(item["case_name"]):
return [doc_id], "case name"
if normalized == self._normal(item["neutral_citation"]):
return [doc_id], "citation"
return [], None
def name_lookup(self, name, k=4):
query_tokens = set(self._normal(name).split()) - {"v", "vs", "versus", "and", "anr", "ors"}
ranked = []
for doc_id, item in self.meta.items():
title_tokens = set(self._normal(item["case_name"]).split())
overlap = len(query_tokens & title_tokens)
if overlap:
ranked.append((overlap, doc_id))
ranked.sort(reverse=True)
return [self._card(doc_id) for _, doc_id in ranked[:k]]
class CaseReferenceResolutionTest(unittest.TestCase):
def setUp(self):
self.corpus = FakeCorpus()
def test_full_party_title_resolves_without_vector_search(self):
result = resolve_case_reference(
self.corpus, "Babu Lal vs Hazari Lal Klshori Lal & Ors"
)
self.assertEqual(result["status"], "resolved")
self.assertEqual(result["case"]["doc_id"], "babu-specific")
def test_short_name_binds_to_active_case(self):
result = resolve_case_reference(
self.corpus, "Babu Lal", active_case_id="babu-specific"
)
self.assertEqual(result["status"], "resolved")
self.assertEqual(result["source"], "active_case")
def test_short_name_binds_to_unique_recent_result(self):
result = resolve_case_reference(
self.corpus, "Babu Lal", recent_case_ids=["babu-specific"]
)
self.assertEqual(result["status"], "resolved")
self.assertEqual(result["source"], "recent_result")
def test_fresh_ambiguous_name_requires_user_selection(self):
result = resolve_case_reference(self.corpus, "Bachan Singh")
self.assertEqual(result["status"], "ambiguous")
self.assertGreaterEqual(len(result["candidates"]), 2)
def test_missing_case_is_reported_as_not_found(self):
result = resolve_case_reference(
self.corpus, "Anne Besant National Girls High School"
)
self.assertEqual(result["status"], "not_found")
self.assertEqual(result["candidates"], [])
if __name__ == "__main__":
unittest.main()