"""
Unit tests for GenerateEmbeddingsOperator.
"""

import pytest
from unittest.mock import Mock, MagicMock, patch, AsyncMock
from plugins.operators.generate_embeddings import GenerateEmbeddingsOperator


class TestGenerateEmbeddingsOperator:
    """Test suite for GenerateEmbeddingsOperator."""

    @pytest.fixture
    def mock_context(self):
        """Create a mock Airflow context."""
        return {'ti': Mock()}

    @pytest.fixture
    def sample_excel_chunks(self):
        """Sample Excel chunk result data."""
        return {
            'chunks': [
                {
                    'content': {'Name': 'Alice', 'Age': '25'},
                    '_source': {'document_id': 'test-doc', 'sheet': 'Sheet1', 'row': 1, 'table_id': 0},
                },
                {
                    'content': {'Name': 'Bob', 'Age': '30'},
                    '_source': {'document_id': 'test-doc', 'sheet': 'Sheet1', 'row': 2, 'table_id': 0},
                },
            ],
            'chunk_count': 2,
            'chunking_strategy': 'table',
        }

    @pytest.fixture
    def sample_text_chunks(self):
        """Sample semantic chunk result data."""
        return {
            'chunks': [
                {
                    'file_id': 'test-doc',
                    'text': 'This is the first chunk of content.',
                    'metadata': {'page_number': 1},
                },
            ],
            'chunk_count': 1,
            'chunking_strategy': 'semantic',
        }

    def test_generate_embeddings_excel_chunks(self, mock_context, sample_excel_chunks):
        """Test embedding generation for Excel table chunks."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'chunk_excel':
                return sample_excel_chunks
            return None

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        with patch('plugins.operators.generate_embeddings.EmbeddingService') as mock_service_class:
            # Mock embedding service
            mock_service = MagicMock()
            mock_embed_batch = AsyncMock(return_value=[
                [0.1] * 3072,  # First embedding
                [0.2] * 3072,  # Second embedding
            ])
            mock_service.embed_batch = mock_embed_batch
            mock_service_class.return_value = mock_service

            op = GenerateEmbeddingsOperator(
                task_id='test_embed',
                batch_size=16,
            )

            result = op.execute(mock_context)

            # Verify result structure
            assert result['embedding_count'] == 2
            assert len(result['embeddings']) == 2
            assert result['model'] == 'text-embedding-3-large'
            assert result['chunks'] == sample_excel_chunks['chunks']

            # Verify embeddings have chunk_id and vector
            for emb in result['embeddings']:
                assert 'chunk_id' in emb
                assert 'vector' in emb
                assert len(emb['vector']) == 3072

    def test_generate_embeddings_text_chunks(self, mock_context, sample_text_chunks):
        """Test embedding generation for semantic text chunks."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'chunk_text':
                return sample_text_chunks
            return None

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        with patch('plugins.operators.generate_embeddings.EmbeddingService') as mock_service_class:
            mock_service = MagicMock()
            mock_embed_batch = AsyncMock(return_value=[
                [0.3] * 3072,
            ])
            mock_service.embed_batch = mock_embed_batch
            mock_service_class.return_value = mock_service

            op = GenerateEmbeddingsOperator(task_id='test_embed')

            result = op.execute(mock_context)

            assert result['embedding_count'] == 1
            assert len(result['embeddings']) == 1

    def test_generate_embeddings_combined_chunks(
        self,
        mock_context,
        sample_excel_chunks,
        sample_text_chunks
    ):
        """Test embedding generation with both Excel and text chunks."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'chunk_excel':
                return sample_excel_chunks
            elif task_ids == 'chunk_text':
                return sample_text_chunks
            return None

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        with patch('plugins.operators.generate_embeddings.EmbeddingService') as mock_service_class:
            mock_service = MagicMock()
            mock_embed_batch = AsyncMock(return_value=[
                [0.1] * 3072,
                [0.2] * 3072,
                [0.3] * 3072,
            ])
            mock_service.embed_batch = mock_embed_batch
            mock_service_class.return_value = mock_service

            op = GenerateEmbeddingsOperator(task_id='test_embed')

            result = op.execute(mock_context)

            # 2 Excel + 1 text chunk = 3 embeddings
            assert result['embedding_count'] == 3
            assert len(result['chunks']) == 3

    def test_generate_embeddings_no_chunks(self, mock_context):
        """Test handling of no chunks."""
        mock_context['ti'].xcom_pull.return_value = None

        op = GenerateEmbeddingsOperator(task_id='test_embed')

        result = op.execute(mock_context)

        assert result['embedding_count'] == 0
        assert result['embeddings'] == []
        assert result['chunks'] == []

    def test_generate_embeddings_empty_chunks(self, mock_context):
        """Test handling of empty chunk lists."""
        def xcom_pull_side_effect(task_ids):
            return {'chunks': [], 'chunk_count': 0}

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        op = GenerateEmbeddingsOperator(task_id='test_embed')

        result = op.execute(mock_context)

        assert result['embedding_count'] == 0
        assert result['embeddings'] == []

    def test_generate_embeddings_batch_processing(self, mock_context):
        """Test that embeddings are processed in batches."""
        # Create 20 chunks (more than default batch size of 16)
        chunks = {
            'chunks': [
                {'content': {'text': f'Content {i}'}, '_source': {'document_id': 'test', 'sheet': 'S', 'row': i, 'table_id': 0}}
                for i in range(20)
            ],
            'chunk_count': 20,
        }

        def xcom_pull_side_effect(task_ids):
            if task_ids == 'chunk_excel':
                return chunks
            return None

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        with patch('plugins.operators.generate_embeddings.EmbeddingService') as mock_service_class:
            mock_service = MagicMock()

            # Track batch sizes
            batch_calls = []

            async def mock_embed(texts):
                batch_calls.append(len(texts))
                return [[0.1] * 3072] * len(texts)

            mock_service.embed_batch = mock_embed
            mock_service_class.return_value = mock_service

            op = GenerateEmbeddingsOperator(
                task_id='test_embed',
                batch_size=16,
            )

            result = op.execute(mock_context)

            assert result['embedding_count'] == 20
            # Should be 2 batches: 16 + 4
            assert len(batch_calls) == 2
            assert batch_calls[0] == 16
            assert batch_calls[1] == 4

    def test_extract_text_from_table_record(self):
        """Test text extraction from table record format."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {
            'content': {'Name': 'Alice', 'Age': '25'},
            '_source': {'document_id': 'test', 'sheet': 'Sheet1', 'row': 1, 'table_id': 0},
        }

        text = op._extract_text_from_chunk(chunk, 0)

        assert 'Name: Alice' in text
        assert 'Age: 25' in text
        assert 'From: test' in text

    def test_extract_text_from_raw_text_record(self):
        """Test text extraction from raw text record format."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {
            'content': {'text': 'This is raw text content'},
            '_source': {'document_id': 'test', 'sheet': 'Sheet1', 'row': -1, 'table_id': 0},
        }

        text = op._extract_text_from_chunk(chunk, 0)

        assert text == 'This is raw text content'

    def test_extract_text_from_semantic_chunk(self):
        """Test text extraction from semantic chunk format."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {
            'text': 'Semantic chunk content',
            'file_id': 'test-file',
            'metadata': {},
        }

        text = op._extract_text_from_chunk(chunk, 0)

        assert text == 'Semantic chunk content'

    def test_extract_chunk_id_from_source(self):
        """Test chunk ID extraction from _source metadata."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {
            'content': {},
            '_source': {'document_id': 'doc123', 'table_id': 2, 'row': 5},
        }

        chunk_id = op._extract_chunk_id(chunk, 0)

        assert chunk_id == 'doc123_t2_r5'

    def test_extract_chunk_id_fallback(self):
        """Test chunk ID fallback when no source available."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {'content': {}}

        chunk_id = op._extract_chunk_id(chunk, 42)

        assert chunk_id == 'chunk_42'

    def test_extract_chunk_id_from_file_id(self):
        """Test chunk ID extraction from file_id field."""
        op = GenerateEmbeddingsOperator(task_id='test')

        chunk = {
            'file_id': 'semantic-chunk-id',
            'text': 'content',
        }

        chunk_id = op._extract_chunk_id(chunk, 0)

        assert chunk_id == 'semantic-chunk-id'

    def test_batch_size_capped_at_api_limit(self):
        """Test that batch size is capped at API limit."""
        op = GenerateEmbeddingsOperator(
            task_id='test',
            batch_size=200,  # Exceeds limit
        )

        assert op.batch_size == GenerateEmbeddingsOperator.MAX_API_BATCH_SIZE
        assert op.batch_size == 100

    def test_custom_task_ids(self):
        """Test using custom task IDs."""
        op = GenerateEmbeddingsOperator(
            task_id='test',
            chunk_excel_task_id='custom_excel',
            chunk_text_task_id='custom_text',
        )

        assert op.chunk_excel_task_id == 'custom_excel'
        assert op.chunk_text_task_id == 'custom_text'

    def test_default_constants(self):
        """Test default constant values."""
        assert GenerateEmbeddingsOperator.DEFAULT_BATCH_SIZE == 16
        assert GenerateEmbeddingsOperator.MAX_API_BATCH_SIZE == 100
        assert GenerateEmbeddingsOperator.DEFAULT_CHUNK_EXCEL_TASK_ID == 'chunk_excel'
        assert GenerateEmbeddingsOperator.DEFAULT_CHUNK_TEXT_TASK_ID == 'chunk_text'
        assert GenerateEmbeddingsOperator.EMBEDDING_MODEL == 'text-embedding-3-large'
        assert GenerateEmbeddingsOperator.EMBEDDING_DIMENSIONS == 3072

    def test_empty_result(self):
        """Test _empty_result helper method."""
        op = GenerateEmbeddingsOperator(task_id='test')

        result = op._empty_result()

        assert result['embeddings'] == []
        assert result['embedding_count'] == 0
        assert result['model'] == 'text-embedding-3-large'
        assert result['chunks'] == []

    def test_skip_empty_text_chunks(self, mock_context):
        """Test that chunks with no text are skipped."""
        chunks = {
            'chunks': [
                {'content': {}, '_source': {}},  # Empty content
                {'content': {'Name': 'Alice'}, '_source': {'document_id': 'test', 'sheet': 'S', 'row': 1, 'table_id': 0}},
            ],
            'chunk_count': 2,
        }

        def xcom_pull_side_effect(task_ids):
            if task_ids == 'chunk_excel':
                return chunks
            return None

        mock_context['ti'].xcom_pull.side_effect = xcom_pull_side_effect

        with patch('plugins.operators.generate_embeddings.EmbeddingService') as mock_service_class:
            mock_service = MagicMock()
            mock_embed_batch = AsyncMock(return_value=[
                [0.1] * 3072,  # Only one embedding (empty chunk skipped)
            ])
            mock_service.embed_batch = mock_embed_batch
            mock_service_class.return_value = mock_service

            op = GenerateEmbeddingsOperator(task_id='test_embed')

            result = op.execute(mock_context)

            # Only one chunk had text, so only one embedding
            assert result['embedding_count'] == 1
