"""Unit tests for semantic retriever module."""

import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import UUID, uuid4

from src.api.schemas.query import SourceReference
from src.core.exceptions import RetrievalError
from src.retrieval.semantic_retriever import SemanticRetriever


class TestSemanticRetriever:
    """Tests for SemanticRetriever class."""

    @pytest.fixture
    def mock_embedding_service(self):
        """Create mock EmbeddingService."""
        service = MagicMock()
        service.embed_batch = AsyncMock(return_value=[[0.1] * 3072])
        return service

    @pytest.fixture
    def mock_milvus_collection(self):
        """Create mock Milvus Collection."""
        collection = MagicMock()
        collection.load = MagicMock()
        collection.search = MagicMock()
        return collection

    @pytest.fixture
    def semantic_retriever(self, mock_embedding_service):
        """Create SemanticRetriever with mocked embedding service."""
        with patch("src.retrieval.semantic_retriever.EmbeddingService") as mock_class:
            mock_class.return_value = mock_embedding_service
            retriever = SemanticRetriever(top_k=5)
            retriever.embedding_service = mock_embedding_service
            return retriever

    @pytest.mark.asyncio
    async def test_retrieve_success(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test successful retrieval returns sources and similarity score."""
        query = "What is item code 2024-4_0019?"

        # Mock search results
        mock_hit = MagicMock()
        mock_hit.distance = 0.85
        mock_hit.entity.get = lambda key, default="": {
            "filename": "distribution_spec.xlsx",
            "sheet_name": "Item Mapping",
            "col_range": "A42:D42",
            "text_content": "Item Code: 2024-4_0019, Report: not applicable",
        }.get(key, default)

        mock_results = [[mock_hit]]
        mock_milvus_collection.search.return_value = mock_results

        with patch("src.retrieval.semantic_retriever.connections") as mock_conn:
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve
                sources, top_similarity = await semantic_retriever.retrieve(query)

                # Verify results
                assert len(sources) == 1
                assert isinstance(sources[0], SourceReference)
                assert sources[0].file == "distribution_spec.xlsx"
                assert sources[0].sheet == "Item Mapping"
                assert sources[0].location == "A42:D42"
                assert "Item Code: 2024-4_0019" in sources[0].context
                assert top_similarity == 0.85

                # Verify embedding was called
                mock_embedding_service.embed_batch.assert_called_once_with([query])

                # Verify Milvus search was called
                mock_milvus_collection.search.assert_called_once()
                search_args = mock_milvus_collection.search.call_args.kwargs
                assert search_args["limit"] == 5
                assert search_args["anns_field"] == "vector"
                assert search_args["param"]["metric_type"] == "COSINE"

                # Verify connection cleanup
                mock_conn.disconnect.assert_called_once_with("default")

    @pytest.mark.asyncio
    async def test_retrieve_with_document_filter(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test retrieval with document_ids filter."""
        query = "test query"
        doc_id_1 = uuid4()
        doc_id_2 = uuid4()
        document_ids = [doc_id_1, doc_id_2]

        # Mock empty results
        mock_milvus_collection.search.return_value = [[]]

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve with document filter
                sources, top_similarity = await semantic_retriever.retrieve(
                    query, document_ids=document_ids
                )

                # Verify filter expression was built correctly
                search_args = mock_milvus_collection.search.call_args.kwargs
                expr = search_args["expr"]
                assert f'"{str(doc_id_1)}"' in expr
                assert f'"{str(doc_id_2)}"' in expr
                assert "document_id in" in expr

    @pytest.mark.asyncio
    async def test_retrieve_empty_results(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test retrieval with no matching results."""
        query = "nonexistent query"

        # Mock empty results
        mock_milvus_collection.search.return_value = [[]]

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve
                sources, top_similarity = await semantic_retriever.retrieve(query)

                # Verify empty results
                assert sources == []
                assert top_similarity == 0.0

    @pytest.mark.asyncio
    async def test_retrieve_multiple_results(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test retrieval returns multiple sources sorted by similarity."""
        query = "test query"

        # Mock multiple search results
        mock_hit_1 = MagicMock()
        mock_hit_1.distance = 0.9
        mock_hit_1.entity.get = lambda key, default="": {
            "filename": "doc1.xlsx",
            "sheet_name": "Sheet1",
            "col_range": "A1:B1",
            "text_content": "Result 1",
        }.get(key, default)

        mock_hit_2 = MagicMock()
        mock_hit_2.distance = 0.75
        mock_hit_2.entity.get = lambda key, default="": {
            "filename": "doc2.xlsx",
            "sheet_name": "Sheet2",
            "col_range": "A2:B2",
            "text_content": "Result 2",
        }.get(key, default)

        mock_results = [[mock_hit_1, mock_hit_2]]
        mock_milvus_collection.search.return_value = mock_results

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve
                sources, top_similarity = await semantic_retriever.retrieve(query)

                # Verify multiple results
                assert len(sources) == 2
                assert sources[0].file == "doc1.xlsx"
                assert sources[1].file == "doc2.xlsx"
                assert top_similarity == 0.9  # Top similarity from first result

    @pytest.mark.asyncio
    async def test_retrieve_connection_error(
        self, semantic_retriever, mock_embedding_service
    ):
        """Test RetrievalError raised on connection failure."""
        query = "test query"

        with patch("src.retrieval.semantic_retriever.connections") as mock_conn:
            mock_conn.connect.side_effect = Exception("Connection failed")

            # Should raise RetrievalError
            with pytest.raises(RetrievalError, match="Failed to retrieve results"):
                await semantic_retriever.retrieve(query)

            # Verify cleanup was attempted
            mock_conn.disconnect.assert_called()

    @pytest.mark.asyncio
    async def test_retrieve_search_error(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test RetrievalError raised on search failure."""
        query = "test query"

        # Mock search error
        mock_milvus_collection.search.side_effect = Exception("Search failed")

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Should raise RetrievalError
                with pytest.raises(RetrievalError, match="Failed to retrieve results"):
                    await semantic_retriever.retrieve(query)

    @pytest.mark.asyncio
    async def test_retrieve_embedding_error(
        self, semantic_retriever, mock_embedding_service
    ):
        """Test RetrievalError raised on embedding generation failure."""
        query = "test query"

        # Mock embedding error
        mock_embedding_service.embed_batch.side_effect = Exception("Embedding failed")

        with patch("src.retrieval.semantic_retriever.connections"):
            # Should raise RetrievalError
            with pytest.raises(RetrievalError, match="Failed to retrieve results"):
                await semantic_retriever.retrieve(query)

    @pytest.mark.asyncio
    async def test_retrieve_custom_top_k(
        self, mock_embedding_service, mock_milvus_collection
    ):
        """Test retrieval with custom top_k parameter."""
        # Create retriever with custom top_k
        with patch("src.retrieval.semantic_retriever.EmbeddingService") as mock_class:
            mock_class.return_value = mock_embedding_service
            retriever = SemanticRetriever(top_k=10)
            retriever.embedding_service = mock_embedding_service

        query = "test query"
        mock_milvus_collection.search.return_value = [[]]

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve
                await retriever.retrieve(query)

                # Verify custom top_k was used
                search_args = mock_milvus_collection.search.call_args.kwargs
                assert search_args["limit"] == 10

    @pytest.mark.asyncio
    async def test_retrieve_missing_fields(
        self, semantic_retriever, mock_embedding_service, mock_milvus_collection
    ):
        """Test retrieval handles missing optional fields gracefully."""
        query = "test query"

        # Mock result with missing fields
        mock_hit = MagicMock()
        mock_hit.distance = 0.8
        mock_hit.entity.get = lambda key, default="": default  # All fields missing

        mock_results = [[mock_hit]]
        mock_milvus_collection.search.return_value = mock_results

        with patch("src.retrieval.semantic_retriever.connections"):
            with patch(
                "src.retrieval.semantic_retriever.Collection",
                return_value=mock_milvus_collection,
            ):
                # Call retrieve
                sources, top_similarity = await semantic_retriever.retrieve(query)

                # Verify empty strings used as defaults
                assert len(sources) == 1
                assert sources[0].file == ""
                assert sources[0].sheet == ""
                assert sources[0].location == ""
                assert sources[0].context == ""
                assert top_similarity == 0.8
