import json import tempfile import unittest from pathlib import Path from unittest.mock import patch from search.evaluate_relevance import evaluate, load_queries, search_body class SearchRelevanceTest(unittest.TestCase): def test_loads_queries_and_calculates_document_metrics(self): queries = [ { "id": "ru-01", "language": "ru", "query": "трудовой договор", "relevant_document_codes": ["7", "8"], }, { "id": "ky-01", "language": "ky", "query": "эмгек келишими", "relevant_document_codes": ["9"], }, ] with tempfile.TemporaryDirectory() as temporary: path = Path(temporary) / "queries.json" path.write_text(json.dumps(queries, ensure_ascii=False), encoding="utf-8") loaded = load_queries(path) responses = [ {"hits": {"hits": [{"_source": {"document_code": code}} for code in ["7", "10", "8"]]}}, {"hits": {"hits": [{"_source": {"document_code": code}} for code in ["10", "9"]]}}, ] with patch("search.evaluate_relevance.request_json", side_effect=responses) as request: result = evaluate(loaded, "http://127.0.0.1:9200", "test", 10) self.assertEqual(result["summary"], {"query_count": 2, "recall_at_10": 1.0, "mrr_at_10": 0.75}) self.assertEqual(result["queries"][0]["reciprocal_rank_at_10"], 1.0) body = json.loads(request.call_args_list[0].args[2]) self.assertFalse(body["track_total_hits"]) self.assertEqual(body["collapse"], {"field": "document_code"}) self.assertEqual(body["query"]["bool"]["must"]["multi_match"]["type"], "cross_fields") def test_company_registration_intent_boosts_current_documents(self): for language, query, status in ( ("ru", "как открыть ОсОО", "Действует"), ("ky", "ЖЧК ачуу тартиби", "Күчүндө"), ): body = json.loads(search_body(language, query, 10)) search_query = body["query"]["bool"] self.assertEqual(search_query["minimum_should_match"], 1) boosts = [clause["constant_score"] for clause in search_query["should"][1:]] self.assertEqual([item["boost"] for item in boosts], [2000, 1000]) self.assertEqual( [item["filter"]["bool"]["filter"][0]["term"]["document_code"] for item in boosts], ["230044970", "667"], ) self.assertTrue(all(item["filter"]["bool"]["filter"][1] == {"term": {f"status_{language}": status}} for item in boosts)) if __name__ == "__main__": unittest.main()