152 lines
5.7 KiB
Python
152 lines
5.7 KiB
Python
"""Measure document search Recall@K and MRR@K against a relevance set."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import json
|
||
import sys
|
||
import urllib.parse
|
||
from pathlib import Path
|
||
|
||
from search.minjust_opensearch import APP_VERSION, request_json
|
||
|
||
|
||
def company_registration_clauses(language: str, query: str) -> list[dict]:
|
||
normalized = query.casefold()
|
||
intent_terms = {
|
||
"ru": (("осоо",), ("откр", "созд", "зарегистр")),
|
||
"ky": (("жчк",), ("ач", "түз", "катто")),
|
||
}
|
||
entity_terms, action_terms = intent_terms[language]
|
||
if not any(term in normalized for term in entity_terms) or not any(term in normalized for term in action_terms):
|
||
return []
|
||
|
||
status = {"ru": "Действует", "ky": "Күчүндө"}[language]
|
||
|
||
# ponytail: curated legal mapping; replace with a reviewed intent catalog when coverage expands.
|
||
def clause(document_code: str, boost: int) -> dict:
|
||
return {
|
||
"constant_score": {
|
||
"filter": {
|
||
"bool": {
|
||
"filter": [
|
||
{"term": {"document_code": document_code}},
|
||
{"term": {f"status_{language}": status}},
|
||
]
|
||
}
|
||
},
|
||
"boost": boost,
|
||
}
|
||
}
|
||
|
||
return [clause("230044970", 2000), clause("667", 1000)]
|
||
|
||
|
||
def search_body(language: str, query: str, top_k: int) -> bytes:
|
||
full_text = {
|
||
"multi_match": {
|
||
"query": query,
|
||
"fields": [f"document_name_{language}", f"text_{language}"],
|
||
"type": "cross_fields",
|
||
}
|
||
}
|
||
clauses = company_registration_clauses(language, query)
|
||
bool_query = {"filter": {"term": {"language": language}}}
|
||
if clauses:
|
||
bool_query.update({"should": [full_text, *clauses], "minimum_should_match": 1})
|
||
else:
|
||
bool_query["must"] = full_text
|
||
return json.dumps({
|
||
"size": top_k,
|
||
"track_total_hits": False,
|
||
"_source": ["document_code"],
|
||
"query": {"bool": bool_query},
|
||
"collapse": {"field": "document_code"},
|
||
}, ensure_ascii=False).encode()
|
||
|
||
|
||
def load_queries(path: Path) -> list[dict]:
|
||
with path.open(encoding="utf-8") as source:
|
||
queries = json.load(source)
|
||
if not isinstance(queries, list) or not queries:
|
||
raise ValueError("Relevance set must be a non-empty JSON array")
|
||
|
||
seen = set()
|
||
for item in queries:
|
||
if not isinstance(item, dict) or set(item) != {
|
||
"id", "language", "query", "relevant_document_codes"
|
||
}:
|
||
raise ValueError("Each query must contain id, language, query and relevant_document_codes")
|
||
codes = item["relevant_document_codes"]
|
||
if (
|
||
not isinstance(item["id"], str)
|
||
or not item["id"].strip()
|
||
or item["id"] in seen
|
||
or not isinstance(item["language"], str)
|
||
or item["language"] not in {"ru", "ky"}
|
||
or not isinstance(item["query"], str)
|
||
or not item["query"].strip()
|
||
or not isinstance(codes, list)
|
||
or not codes
|
||
or any(not isinstance(code, str) or not code for code in codes)
|
||
or len(codes) != len(set(codes))
|
||
):
|
||
raise ValueError(f"Invalid relevance query: {item.get('id', '<unknown>')}")
|
||
seen.add(item["id"])
|
||
return queries
|
||
|
||
|
||
def search(base_url: str, index: str, item: dict, top_k: int) -> list[str]:
|
||
language = item["language"]
|
||
body = search_body(language, item["query"], top_k)
|
||
url = f"{base_url.rstrip('/')}/{urllib.parse.quote(index, safe='')}/_search"
|
||
response = request_json(url, "POST", body, "application/json")
|
||
try:
|
||
return [hit["_source"]["document_code"] for hit in response["hits"]["hits"]]
|
||
except (KeyError, TypeError) as error:
|
||
raise RuntimeError("OpenSearch search response is incomplete") from error
|
||
|
||
|
||
def evaluate(queries: list[dict], base_url: str, index: str, top_k: int) -> dict:
|
||
results = []
|
||
for item in queries:
|
||
retrieved = search(base_url, index, item, top_k)
|
||
relevant = set(item["relevant_document_codes"])
|
||
matches = [rank for rank, code in enumerate(retrieved, 1) if code in relevant]
|
||
results.append({
|
||
"id": item["id"],
|
||
"language": item["language"],
|
||
"query": item["query"],
|
||
"retrieved_document_codes": retrieved,
|
||
f"recall_at_{top_k}": len(relevant.intersection(retrieved)) / len(relevant),
|
||
f"reciprocal_rank_at_{top_k}": 1 / matches[0] if matches else 0.0,
|
||
})
|
||
return {
|
||
"summary": {
|
||
"query_count": len(results),
|
||
f"recall_at_{top_k}": sum(item[f"recall_at_{top_k}"] for item in results) / len(results),
|
||
f"mrr_at_{top_k}": sum(item[f"reciprocal_rank_at_{top_k}"] for item in results) / len(results),
|
||
},
|
||
"queries": results,
|
||
}
|
||
|
||
|
||
def main() -> int:
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument("relevance_set", type=Path)
|
||
parser.add_argument("--url", default="http://127.0.0.1:9200")
|
||
parser.add_argument("--index", default="akyldash-fragments-v1")
|
||
parser.add_argument("--top-k", type=int, default=10)
|
||
parser.add_argument("--version", action="version", version=APP_VERSION)
|
||
arguments = parser.parse_args()
|
||
if arguments.top_k <= 0:
|
||
raise SystemExit("--top-k must be greater than zero")
|
||
result = evaluate(load_queries(arguments.relevance_set), arguments.url, arguments.index, arguments.top_k)
|
||
print(json.dumps(result, ensure_ascii=False, indent=2))
|
||
print(f"Akyldash Backend v{APP_VERSION} · Frontend — not created", file=sys.stderr)
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
raise SystemExit(main())
|