"""
Domain-Aware Retriever for optimized RAG queries.

Implements two-phase retrieval:
1. Pre-filter by domain terms and category
2. Vector similarity search on filtered candidates
"""

import time
from typing import Any, Optional

import structlog
from pymilvus import Collection, connections

from src.core.config import settings
from src.knowledge.embeddings import EmbeddingService
from src.knowledge.query_term_matcher import QueryTermMatcher

logger = structlog.get_logger(__name__)


class DomainAwareRetriever:
    """Retriever with domain-based pre-filtering for faster queries.

    Two-phase retrieval strategy:
    1. Filter by domain_category and domain_terms (reduces candidates 10-100x)
    2. Vector similarity search on filtered subset

    Falls back to full search if no domain terms match.
    """

    def __init__(
        self,
        query_matcher: Optional[QueryTermMatcher] = None,
        embedding_service: Optional[EmbeddingService] = None,
        collection_name: Optional[str] = None,
    ):
        """Initialize domain-aware retriever.

        Args:
            query_matcher: QueryTermMatcher for term extraction
            embedding_service: EmbeddingService for query embedding
            collection_name: Milvus collection name (uses settings default if None)
        """
        self.query_matcher = query_matcher or QueryTermMatcher()
        self.embedding_service = embedding_service or EmbeddingService()
        self.collection_name = collection_name or settings.milvus_collection
        self.logger = structlog.get_logger(__name__)
        self._connected = False

    def _connect(self) -> None:
        """Connect to Milvus if not already connected."""
        if not self._connected:
            connections.connect(
                alias="default",
                host=settings.milvus_host,
                port=settings.milvus_port,
            )
            self._connected = True

    def _disconnect(self) -> None:
        """Disconnect from Milvus."""
        if self._connected:
            connections.disconnect("default")
            self._connected = False

    async def retrieve(
        self,
        query: str,
        tenant_id: str,
        node_id: str,
        top_k: int = 10,
        document_id: Optional[str] = None,
        tag_filters: Optional[dict[str, list[str]]] = None,
    ) -> list[dict[str, Any]]:
        """Retrieve relevant chunks using domain-aware filtering.

        Args:
            query: User query string
            tenant_id: Tenant UUID string (REQUIRED for data isolation)
            node_id: Node UUID string (REQUIRED for data isolation)
            top_k: Number of results to return
            document_id: Optional document ID filter
            tag_filters: Optional tag filtering with operators:
                - "contains_any": Match records with ANY of these tags
                - "contains_all": Match records with ALL of these tags
                Example: {"contains_any": ["req:functional"], "contains_all": ["topic:auth"]}

        Returns:
            List of result dicts with fields:
                - record_id: PostgreSQL record UUID
                - document_id: Document UUID
                - filename: Source filename
                - text_content: Chunk text
                - metadata: JSON metadata
                - domain_category: Domain category
                - domain_terms: List of domain terms
                - tags: Semantic tags
                - score: Similarity score

        Example:
            >>> retriever = DomainAwareRetriever()
            >>> results = await retriever.retrieve("What is Q4 EBITDA?")
            >>> results[0]["domain_category"]
            "finance"
            >>> # With tag filtering
            >>> results = await retriever.retrieve(
            ...     "requirements",
            ...     tag_filters={"contains_any": ["req:functional", "req:interface"]}
            ... )
        """
        start_time = time.perf_counter()

        # Extract domain terms from query
        query_analysis = await self.query_matcher.extract_query_terms(
            query, tenant_id
        )

        matched_terms = query_analysis["matched_terms"]
        domain_category = query_analysis["domain_category"]

        self.logger.info(
            "query_analysis",
            matched_terms=matched_terms,
            category=domain_category,
            matching_method=query_analysis["matching_method"]
        )

        # Generate query embedding
        query_vectors = await self.embedding_service.embed_batch([query])
        query_vector = query_vectors[0]

        # Build filter expression
        filter_expr = self._build_filter(
            matched_terms,
            domain_category,
            tenant_id,
            node_id,
            document_id,
            tag_filters,
        )

        # Execute filtered search
        self._connect()
        collection = Collection(self.collection_name)
        collection.load()

        try:
            # Search parameters
            search_params = {
                "metric_type": "COSINE",
                "params": {"ef": 128},  # HNSW search parameter
            }

            # Perform vector search with filter
            search_results = collection.search(
                data=[query_vector],
                anns_field="vector",
                param=search_params,
                limit=top_k,
                expr=filter_expr if filter_expr else None,
                output_fields=[
                    "record_id",
                    "document_id",
                    "filename",
                    "tags",
                    "document_type",
                    "text_content",
                    "metadata",
                    "domain_category",
                    "domain_terms",
                ],
            )

            # Format results
            results = []
            for hits in search_results:
                for hit in hits:
                    result = {
                        "record_id": hit.entity.get("record_id"),
                        "document_id": hit.entity.get("document_id"),
                        "filename": hit.entity.get("filename"),
                        "document_type": hit.entity.get("document_type"),
                        "text_content": hit.entity.get("text_content"),
                        "metadata": hit.entity.get("metadata"),
                        "domain_category": hit.entity.get("domain_category"),
                        "domain_terms": hit.entity.get("domain_terms"),
                        "score": hit.score,
                    }
                    results.append(result)

            latency_ms = (time.perf_counter() - start_time) * 1000

            self.logger.info(
                "retrieval_complete",
                results_count=len(results),
                latency_ms=round(latency_ms, 2),
                filter_used=bool(filter_expr),
                matched_terms_count=len(matched_terms),
            )

            return results

        except Exception as e:
            self.logger.error(
                "retrieval_failed",
                error=str(e),
                query=query[:100],
            )
            raise

    def _build_filter(
        self,
        matched_terms: list[str],
        domain_category: str,
        tenant_id: str,
        node_id: str,
        document_id: Optional[str] = None,
        tag_filters: Optional[dict[str, list[str]]] = None,
    ) -> Optional[str]:
        """Build Milvus filter expression for domain pre-filtering.

        Implements multi-level filtering strategy:
        1. tenant_id + node_id (MANDATORY — data isolation)
        2. domain_category (coarse filter)
        3. domain_terms (fine-grained filter using array_contains_any)
        4. semantic tags (contains_any / contains_all)

        Args:
            matched_terms: List of matched domain terms
            domain_category: Inferred domain category
            tenant_id: Tenant UUID string (REQUIRED)
            node_id: Node UUID string (REQUIRED)
            document_id: Optional document ID filter
            tag_filters: Optional tag filtering with operators

        Returns:
            Milvus filter expression string (always non-None due to mandatory tenant filter)
        """
        filters = []

        # MANDATORY — tenant + node isolation (data-isolation-spec Section 3.3)
        filters.append(f'tenant_id == "{tenant_id}"')
        filters.append(f'node_id == "{node_id}"')

        # Level 1: Document ID filter (if specified)
        if document_id:
            filters.append(f'document_id == "{document_id}"')

        # Level 2: Domain category filter (if not general)
        if domain_category and domain_category != "general":
            filters.append(f'domain_category == "{domain_category}"')

        # Level 3: Domain terms filter (if terms matched)
        if matched_terms:
            # Use array_contains_any for efficient ARRAY field filtering
            # Escape single quotes in terms
            escaped_terms = [term.replace("'", "\\'") for term in matched_terms]
            terms_list = ", ".join(f'"{term}"' for term in escaped_terms)
            filters.append(f"array_contains_any(domain_terms, [{terms_list}])")

        # Level 4: Tag filters (if specified)
        if tag_filters:
            # contains_any: Match records with ANY of these tags
            if "contains_any" in tag_filters and tag_filters["contains_any"]:
                tags = tag_filters["contains_any"]
                escaped_tags = [tag.replace("'", "\\'") for tag in tags]
                tags_list = ", ".join(f'"{tag}"' for tag in escaped_tags)
                filters.append(f"array_contains_any(tags, [{tags_list}])")

            # contains_all: Match records with ALL of these tags
            if "contains_all" in tag_filters and tag_filters["contains_all"]:
                tags = tag_filters["contains_all"]
                # For contains_all, use array_contains for each tag and AND them
                for tag in tags:
                    escaped_tag = tag.replace("'", "\\'")
                    filters.append(f'array_contains(tags, "{escaped_tag}")')

        # Combine filters with AND (always non-empty due to mandatory tenant/node)
        filter_expr = " and ".join(filters)
        self.logger.debug("filter_built", expression=filter_expr)
        return filter_expr

    def __enter__(self):
        """Context manager entry."""
        self._connect()
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):
        """Context manager exit."""
        self._disconnect()
