105 lines
6.3 KiB
Python
105 lines
6.3 KiB
Python
import json
|
||
import tempfile
|
||
import unittest
|
||
from pathlib import Path
|
||
from unittest.mock import patch
|
||
|
||
from search.evaluate_relevance import evaluate, load_queries
|
||
from search.query import COMPANY_REGISTRATION_DOCUMENTS, TAX_PROFIT_DOCUMENTS, build_search_body, query_variants
|
||
|
||
|
||
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)
|
||
|
||
with patch("search.evaluate_relevance.search_documents", side_effect=[["7", "10", "8"], ["10", "9"]]):
|
||
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(build_search_body("ru", "трудовой договор", 10))
|
||
self.assertFalse(body["track_total_hits"])
|
||
self.assertEqual(body["collapse"], {"field": "document_code"})
|
||
self.assertIn({"term": {"is_current_edition": True}}, body["query"]["bool"]["filter"])
|
||
lexical_query = body["query"]["bool"]["must"]["bool"]["should"][0]["match"]["text_ru"]
|
||
self.assertEqual(lexical_query["operator"], "AND")
|
||
self.assertTrue(any("match_phrase" in clause for clause in body["query"]["bool"]["must"]["bool"]["should"]))
|
||
|
||
def test_legal_form_abbreviations_expand_in_both_directions(self):
|
||
self.assertIn("как открыть общество с ограниченной ответственностью", query_variants("ru", "как открыть ОсОО"))
|
||
self.assertIn("регистрация осоо", [variant.casefold() for variant in query_variants("ru", "регистрация общества с ограниченной ответственностью")])
|
||
self.assertIn("жчк ачуу тартиби", [variant.casefold() for variant in query_variants("ky", "жоопкерчилиги чектелген коом ачуу тартиби")])
|
||
|
||
def test_company_registration_intent_boosts_current_documents(self):
|
||
for language, query, status in (
|
||
("ru", "как открыть ОсОО", "Действует"),
|
||
("ky", "ЖЧК ачуу тартиби", "Күчүндө"),
|
||
):
|
||
body = json.loads(build_search_body(language, query, 10))
|
||
search_query = body["query"]["bool"]["must"]["bool"]
|
||
self.assertEqual(search_query["minimum_should_match"], 1)
|
||
self.assertNotIn("filter", search_query)
|
||
lexical = [clause["match"] for clause in search_query["should"] if "match" in clause]
|
||
self.assertTrue(all(next(iter(clause.values()))["operator"] == "OR" for clause in lexical))
|
||
boosts = [clause["constant_score"] for clause in search_query["should"] if "constant_score" in clause]
|
||
self.assertEqual([item["boost"] for item in boosts], [boost for _, boost in COMPANY_REGISTRATION_DOCUMENTS])
|
||
self.assertEqual(
|
||
[item["filter"]["bool"]["filter"][0]["term"]["document_code"] for item in boosts],
|
||
[code for code, _ in COMPANY_REGISTRATION_DOCUMENTS],
|
||
)
|
||
self.assertTrue(all(item["filter"]["bool"]["filter"][1] == {"term": {f"status_{language}": status}} for item in boosts))
|
||
|
||
def test_profit_tax_intent_prioritizes_the_current_tax_code(self):
|
||
for query in ("налог на прибыль", "налога на прибыль организаций", "налогообложение прибыли"):
|
||
body = json.loads(build_search_body("ru", query, 10))
|
||
clauses = body["query"]["bool"]["must"]["bool"]["should"]
|
||
boosts = [clause["constant_score"] for clause in clauses if "constant_score" in clause]
|
||
self.assertEqual([item["boost"] for item in boosts], [boost for _, boost in TAX_PROFIT_DOCUMENTS])
|
||
self.assertEqual([item["filter"]["bool"]["filter"][0]["term"]["document_code"] for item in boosts], ["112340"])
|
||
self.assertTrue(all(item["filter"]["bool"]["filter"][1] == {"term": {"status_ru": "Действует"}} for item in boosts))
|
||
|
||
for query in ("прибыль организации", "налог на доходы", "налог на имущество"):
|
||
body = json.loads(build_search_body("ru", query, 10))
|
||
clauses = body["query"]["bool"]["must"]["bool"]["should"]
|
||
self.assertFalse(any("constant_score" in clause for clause in clauses))
|
||
|
||
def test_company_registration_intent_ignores_non_procedural_queries(self):
|
||
for language, query in (
|
||
("ru", "ОсОО зарегистрирован?"),
|
||
("ru", "кто зарегистрировал ОсОО"),
|
||
("ru", "как открыть счет ОсОО"),
|
||
("ru", "как открыть ОсОО банковский счет"),
|
||
("ru", "как открыть ОсОО банковский счёт"),
|
||
("ru", "как открыть филиал ОсОО"),
|
||
("ru", "как создать договор для ОсОО"),
|
||
("ru", "порядок создания логотипа ОсОО"),
|
||
("ky", "ЖЧК ачык маалымат"),
|
||
("ky", "ЖЧК кантип банк эсебин ачуу"),
|
||
("ky", "ЖЧК кантип келишим түзүү"),
|
||
("ky", "ЖЧК кантип логотип түзүү"),
|
||
):
|
||
body = json.loads(build_search_body(language, query, 10))
|
||
clauses = body["query"]["bool"]["must"]["bool"]["should"]
|
||
self.assertFalse(any("constant_score" in clause for clause in clauses))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|