"""Unit tests for answer generator module."""

import pytest
from unittest.mock import AsyncMock, MagicMock, patch

from openai.types.chat import ChatCompletion, ChatCompletionMessage
from openai.types.chat.chat_completion import Choice

from src.api.schemas.query import SourceReference
from src.core.exceptions import QueryError
from src.retrieval.answer_generator import AnswerGenerator, calculate_confidence


class TestCalculateConfidence:
    """Tests for calculate_confidence function."""

    def test_no_sources_returns_zero(self):
        """Test confidence is 0.0 when no sources."""
        confidence = calculate_confidence([], top_similarity=0.5)
        assert confidence == 0.0

    def test_low_similarity_range(self):
        """Test confidence in 0.3-0.5 range for low similarity (<0.3)."""
        sources = [MagicMock(spec=SourceReference)]

        # Test bottom of range
        confidence = calculate_confidence(sources, top_similarity=0.0)
        assert 0.3 <= confidence <= 0.5

        # Test middle of range
        confidence = calculate_confidence(sources, top_similarity=0.15)
        assert 0.3 <= confidence <= 0.5

        # Test top of range
        confidence = calculate_confidence(sources, top_similarity=0.29)
        assert 0.3 <= confidence <= 0.5

    def test_medium_similarity_range(self):
        """Test confidence in 0.5-0.7 range for medium similarity (0.3-0.6)."""
        sources = [MagicMock(spec=SourceReference)]

        # Test bottom of range
        confidence = calculate_confidence(sources, top_similarity=0.3)
        assert 0.5 <= confidence <= 0.7

        # Test middle of range
        confidence = calculate_confidence(sources, top_similarity=0.45)
        assert 0.5 <= confidence <= 0.7

        # Test top of range
        confidence = calculate_confidence(sources, top_similarity=0.59)
        assert 0.5 <= confidence <= 0.7

    def test_high_similarity_range(self):
        """Test confidence in 0.7-1.0 range for high similarity (>0.6)."""
        sources = [MagicMock(spec=SourceReference)]

        # Test bottom of range
        confidence = calculate_confidence(sources, top_similarity=0.6)
        assert 0.7 <= confidence <= 1.0

        # Test middle of range
        confidence = calculate_confidence(sources, top_similarity=0.8)
        assert 0.7 <= confidence <= 1.0

        # Test top of range (perfect match)
        confidence = calculate_confidence(sources, top_similarity=1.0)
        assert confidence == 1.0

    def test_confidence_increases_with_similarity(self):
        """Test confidence monotonically increases with similarity."""
        sources = [MagicMock(spec=SourceReference)]

        conf_0_1 = calculate_confidence(sources, top_similarity=0.1)
        conf_0_4 = calculate_confidence(sources, top_similarity=0.4)
        conf_0_7 = calculate_confidence(sources, top_similarity=0.7)
        conf_0_9 = calculate_confidence(sources, top_similarity=0.9)

        assert conf_0_1 < conf_0_4 < conf_0_7 < conf_0_9


class TestAnswerGenerator:
    """Tests for AnswerGenerator class."""

    @pytest.fixture
    def mock_openai_client(self):
        """Create mock Azure OpenAI client."""
        client = MagicMock()
        client.chat = MagicMock()
        client.chat.completions = MagicMock()
        return client

    @pytest.fixture
    def answer_generator(self, mock_openai_client):
        """Create AnswerGenerator with mocked client."""
        with patch("src.retrieval.answer_generator.AsyncAzureOpenAI") as mock_class:
            mock_class.return_value = mock_openai_client
            generator = AnswerGenerator()
            generator.client = mock_openai_client
            return generator

    @pytest.fixture
    def sample_sources(self):
        """Create sample source references."""
        return [
            SourceReference(
                file="distribution_spec.xlsx",
                sheet="Item Mapping",
                location="A42:D42",
                context="Item Code: 2024-4_0019, Report: not applicable",
            ),
            SourceReference(
                file="distribution_spec.xlsx",
                sheet="Item Mapping",
                location="A43:D43",
                context="Item Code: 2024-4_0020, Report: applicable",
            ),
        ]

    @pytest.mark.asyncio
    async def test_generate_success(
        self, answer_generator, mock_openai_client, sample_sources
    ):
        """Test successful answer generation."""
        query = "What is item code 2024-4_0019?"

        # Mock LLM response
        mock_message = ChatCompletionMessage(
            role="assistant",
            content="Item code 2024-4_0019 appears in the Item Mapping sheet.",
        )
        mock_choice = Choice(
            index=0,
            message=mock_message,
            finish_reason="stop",
        )
        mock_completion = ChatCompletion(
            id="test-completion",
            model="gpt-4",
            object="chat.completion",
            created=1234567890,
            choices=[mock_choice],
        )
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )

        # Call generate
        answer, confidence = await answer_generator.generate(
            query=query, sources=sample_sources, top_similarity=0.85
        )

        # Verify answer and confidence
        assert "Item code 2024-4_0019" in answer
        assert 0.7 <= confidence <= 1.0  # High similarity -> high confidence

        # Verify LLM was called correctly
        mock_openai_client.chat.completions.create.assert_called_once()
        call_args = mock_openai_client.chat.completions.create.call_args.kwargs
        assert call_args["temperature"] == 0.3
        assert call_args["max_tokens"] == 500
        assert len(call_args["messages"]) == 2
        assert call_args["messages"][0]["role"] == "system"
        assert call_args["messages"][1]["role"] == "user"
        assert query in call_args["messages"][1]["content"]

    @pytest.mark.asyncio
    async def test_generate_no_sources(self, answer_generator, mock_openai_client):
        """Test answer generation with no sources."""
        query = "What is the meaning of life?"

        # Call generate with empty sources
        answer, confidence = await answer_generator.generate(
            query=query, sources=[], top_similarity=0.0
        )

        # Verify no-sources response
        assert "couldn't find any information" in answer.lower()
        assert confidence == 0.1

        # Verify LLM was NOT called
        mock_openai_client.chat.completions.create.assert_not_called()

    @pytest.mark.asyncio
    async def test_generate_formats_context_correctly(
        self, answer_generator, mock_openai_client, sample_sources
    ):
        """Test that context is formatted correctly from sources."""
        query = "test query"

        # Mock LLM response
        mock_message = ChatCompletionMessage(role="assistant", content="Test answer")
        mock_choice = Choice(
            index=0, message=mock_message, finish_reason="stop"
        )
        mock_completion = ChatCompletion(
            id="test", model="gpt-4", object="chat.completion",
            created=1234567890, choices=[mock_choice]
        )
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )

        # Call generate
        await answer_generator.generate(
            query=query, sources=sample_sources, top_similarity=0.7
        )

        # Verify context formatting
        call_args = mock_openai_client.chat.completions.create.call_args.kwargs
        user_prompt = call_args["messages"][1]["content"]

        # Check source numbering and formatting
        assert "Source 1" in user_prompt
        assert "Source 2" in user_prompt
        assert "distribution_spec.xlsx" in user_prompt
        assert "Item Mapping" in user_prompt
        assert "A42:D42" in user_prompt
        assert "2024-4_0019" in user_prompt

    @pytest.mark.asyncio
    async def test_generate_retry_on_failure(
        self, answer_generator, mock_openai_client
    ):
        """Test retry logic on transient LLM errors."""
        query = "test query"
        sources = [
            SourceReference(
                file="test.xlsx", sheet="Sheet1", location="A1", context="test"
            )
        ]

        # Mock LLM to fail twice then succeed
        mock_message = ChatCompletionMessage(role="assistant", content="Success!")
        mock_choice = Choice(index=0, message=mock_message, finish_reason="stop")
        mock_completion = ChatCompletion(
            id="test", model="gpt-4", object="chat.completion",
            created=1234567890, choices=[mock_choice]
        )

        mock_openai_client.chat.completions.create = AsyncMock(
            side_effect=[
                Exception("Rate limit exceeded"),
                Exception("Timeout"),
                mock_completion,
            ]
        )

        # Call generate (should succeed on 3rd attempt)
        answer, confidence = await answer_generator.generate(
            query=query, sources=sources, top_similarity=0.7
        )

        # Verify success after retries
        assert answer == "Success!"
        assert mock_openai_client.chat.completions.create.call_count == 3

    @pytest.mark.asyncio
    async def test_generate_permanent_failure(
        self, answer_generator, mock_openai_client
    ):
        """Test QueryError raised after max retries."""
        query = "test query"
        sources = [
            SourceReference(
                file="test.xlsx", sheet="Sheet1", location="A1", context="test"
            )
        ]

        # Mock LLM to always fail
        mock_openai_client.chat.completions.create = AsyncMock(
            side_effect=Exception("Permanent error")
        )

        # Should raise QueryError after 3 retries
        with pytest.raises(QueryError, match="Failed to generate answer"):
            await answer_generator.generate(
                query=query, sources=sources, top_similarity=0.7
            )

        # Verify 3 attempts were made
        assert mock_openai_client.chat.completions.create.call_count == 3

    @pytest.mark.asyncio
    async def test_generate_empty_llm_response(
        self, answer_generator, mock_openai_client
    ):
        """Test QueryError raised when LLM returns empty content."""
        query = "test query"
        sources = [
            SourceReference(
                file="test.xlsx", sheet="Sheet1", location="A1", context="test"
            )
        ]

        # Mock LLM to return empty content
        mock_message = ChatCompletionMessage(role="assistant", content=None)
        mock_choice = Choice(index=0, message=mock_message, finish_reason="stop")
        mock_completion = ChatCompletion(
            id="test", model="gpt-4", object="chat.completion",
            created=1234567890, choices=[mock_choice]
        )
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )

        # Should raise QueryError and retry
        with pytest.raises(QueryError, match="LLM returned empty response"):
            await answer_generator.generate(
                query=query, sources=sources, top_similarity=0.7
            )

    @pytest.mark.asyncio
    async def test_generate_confidence_varies_with_similarity(
        self, answer_generator, mock_openai_client
    ):
        """Test confidence score varies based on top_similarity."""
        query = "test query"
        sources = [
            SourceReference(
                file="test.xlsx", sheet="Sheet1", location="A1", context="test"
            )
        ]

        # Mock LLM response
        mock_message = ChatCompletionMessage(role="assistant", content="Test answer")
        mock_choice = Choice(index=0, message=mock_message, finish_reason="stop")
        mock_completion = ChatCompletion(
            id="test", model="gpt-4", object="chat.completion",
            created=1234567890, choices=[mock_choice]
        )
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )

        # Test with low similarity
        _, conf_low = await answer_generator.generate(
            query=query, sources=sources, top_similarity=0.2
        )

        # Test with high similarity
        mock_openai_client.chat.completions.create.reset_mock()
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )
        _, conf_high = await answer_generator.generate(
            query=query, sources=sources, top_similarity=0.9
        )

        # High similarity should yield higher confidence
        assert conf_high > conf_low
        assert 0.3 <= conf_low <= 0.5
        assert 0.7 <= conf_high <= 1.0

    @pytest.mark.asyncio
    async def test_generate_system_prompt_content(
        self, answer_generator, mock_openai_client, sample_sources
    ):
        """Test system prompt instructs to cite sources and be concise."""
        query = "test query"

        # Mock LLM response
        mock_message = ChatCompletionMessage(role="assistant", content="Test answer")
        mock_choice = Choice(index=0, message=mock_message, finish_reason="stop")
        mock_completion = ChatCompletion(
            id="test", model="gpt-4", object="chat.completion",
            created=1234567890, choices=[mock_choice]
        )
        mock_openai_client.chat.completions.create = AsyncMock(
            return_value=mock_completion
        )

        # Call generate
        await answer_generator.generate(
            query=query, sources=sample_sources, top_similarity=0.7
        )

        # Verify system prompt content
        call_args = mock_openai_client.chat.completions.create.call_args.kwargs
        system_prompt = call_args["messages"][0]["content"]

        assert "cite your sources" in system_prompt.lower()
        assert "based only on the context" in system_prompt.lower()
        assert "helpful assistant" in system_prompt.lower()
