"""Hybrid Search Service combining BM25 Lexical Search and Dense Vector Search (Qdrant) with Reciprocal Rank Fusion (RRF)."""

from __future__ import annotations

import math
import re
import unicodedata
from collections import Counter
from typing import Any

from app.modules.core.application.embeddings import EmbeddingService
from app.modules.core.application.qdrant_store import QdrantStore


def strip_accents(text: str) -> str:
    """Loại bỏ dấu tiếng Việt để đối sánh linh hoạt."""
    return "".join(
        ch
        for ch in unicodedata.normalize("NFKD", text)
        if not unicodedata.combining(ch)
    )


def bm25_tokenize(text: str) -> list[str]:
    """Phân tách từ chuẩn hóa cho BM25 (hỗ trợ tiếng Việt và ký tự đặc biệt như số hiệu văn bản/mã vật tư)."""
    if not text:
        return []
    normalized = text.lower()
    # Tách các cụm số hiệu như 135/2025/qh15, 206/2026/nd-cp, 12/2021/tt-bxd
    tokens = re.findall(
        r"[a-z0-9àáảãạăằắẳẵặâầấẩẫậèéẻẽẹêềếểễệìíỉĩịòóỏõọôồốổỗộơờớởỡợùúủũụưừứửữựỳýỷỹỵđ/_.\-]+",
        normalized,
    )
    # Thêm cả dạng không dấu để tăng recall
    unaccented_tokens = [strip_accents(t) for t in tokens if strip_accents(t) != t]
    return tokens + unaccented_tokens


class HybridSearchService:
    """Dịch vụ tìm kiếm hỗn hợp (BM25 Lexical + Dense Vectors Qdrant + RRF Reranking)."""

    def __init__(
        self, qdrant_store: QdrantStore, embedding_service: EmbeddingService
    ) -> None:
        self._qdrant_store = qdrant_store
        self._embedding_service = embedding_service
        self._k1 = 1.5
        self._b = 0.75

    def _compute_bm25_score(
        self,
        query_tokens: list[str],
        doc_tokens: list[str],
        avg_dl: float,
        doc_len: int,
        term_idf: dict[str, float],
    ) -> float:
        """Tính điểm BM25 cho một tài liệu đối với câu truy vấn."""
        if not doc_tokens or not query_tokens:
            return 0.0

        doc_counter = Counter(doc_tokens)
        score = 0.0

        for q_term in query_tokens:
            if q_term not in doc_counter:
                continue
            tf = doc_counter[q_term]
            idf = term_idf.get(q_term, 1.0)
            numerator = tf * (self._k1 + 1.0)
            denominator = tf + self._k1 * (
                1.0 - self._b + self._b * (doc_len / max(1.0, avg_dl))
            )
            score += idf * (numerator / max(1e-6, denominator))

        return score

    async def search(
        self,
        query: str,
        top_k: int = 5,
        task_type: str | None = None,
        retrieval_filters: dict[str, Any] | None = None,
        mode: str = "hybrid",  # 'hybrid', 'dense', 'bm25'
    ) -> list[dict[str, Any]]:
        """Tìm kiếm ngữ nghĩa và từ khóa kết hợp."""
        if not query or not query.strip():
            return []

        # 1. Thu thập ứng viên từ Dense Vector Search (Lấy top_k * 4 để rerank)
        candidate_limit = max(top_k * 4, 20)
        query_vector = self._embedding_service.embed_query(query)

        dense_results: list[dict[str, Any]] = []
        try:
            dense_results = self._qdrant_store.search(
                query_vector=query_vector,
                limit=candidate_limit,
                filters=retrieval_filters,
            )
        except Exception:
            dense_results = []

        if mode == "dense":
            return dense_results[:top_k]

        if not dense_results:
            return []

        # 2. Xây dựng chỉ mục BM25 trên tập ứng viên
        query_tokens = bm25_tokenize(query)
        doc_token_list = [
            bm25_tokenize(res.get("text", "") + " " + str(res.get("metadata", {})))
            for res in dense_results
        ]
        doc_lengths = [len(tokens) for tokens in doc_token_list]
        avg_dl = sum(doc_lengths) / max(1, len(doc_lengths))
        n_docs = len(dense_results)

        # Tính IDF cho từng term trong query
        term_idf: dict[str, float] = {}
        for q_term in set(query_tokens):
            n_containing = sum(1 for tokens in doc_token_list if q_term in tokens)
            term_idf[q_term] = math.log(
                1.0 + (n_docs - n_containing + 0.5) / (n_containing + 0.5)
            )

        # 3. Tính điểm BM25 cho từng tài liệu
        bm25_scored_items: list[tuple[int, float]] = []
        for idx, (doc_tokens, doc_len) in enumerate(zip(doc_token_list, doc_lengths)):
            score = self._compute_bm25_score(
                query_tokens, doc_tokens, avg_dl, doc_len, term_idf
            )
            bm25_scored_items.append((idx, score))

        if mode == "bm25":
            bm25_scored_items.sort(key=lambda item: item[1], reverse=True)
            results = []
            for rank, (idx, score) in enumerate(bm25_scored_items[:top_k], start=1):
                item = dict(dense_results[idx])
                item["bm25_score"] = score
                item["rank"] = rank
                results.append(item)
            return results

        # 4. Hợp nhất thứ hạng Reciprocal Rank Fusion (RRF)
        # RRF_Score = 0.6 / (60 + Dense_Rank) + 0.4 / (60 + BM25_Rank)
        dense_ranks = {
            res.get("id", str(i)): i + 1 for i, res in enumerate(dense_results)
        }

        # Sort by BM25
        sorted_by_bm25 = sorted(
            bm25_scored_items, key=lambda item: item[1], reverse=True
        )
        bm25_ranks = {
            dense_results[idx].get("id", str(idx)): rank + 1
            for rank, (idx, _) in enumerate(sorted_by_bm25)
        }
        bm25_score_map = {
            dense_results[idx].get("id", str(idx)): score
            for idx, score in bm25_scored_items
        }

        k = 60.0
        w_dense = 0.6
        w_bm25 = 0.4

        fused_items: list[dict[str, Any]] = []
        for res in dense_results:
            doc_id = res.get("id", "")
            d_rank = dense_ranks.get(doc_id, len(dense_results))
            b_score = bm25_score_map.get(doc_id, 0.0)

            dense_component = w_dense / (k + d_rank)
            bm25_component = (
                (w_bm25 / (k + bm25_ranks.get(doc_id, len(dense_results))))
                if b_score > 0
                else 0.0
            )

            rrf_score = dense_component + bm25_component

            item = dict(res)
            item["rrf_score"] = round(rrf_score, 6)
            item["dense_score"] = res.get("score", 0.0)
            item["bm25_score"] = round(b_score, 4)
            item["search_mode"] = "hybrid_rrf"
            fused_items.append(item)

        fused_items.sort(key=lambda x: x.get("rrf_score", 0.0), reverse=True)
        return fused_items[:top_k]
