"""
Unit tests for StoreChunksOperator.
"""

import pytest
from unittest.mock import Mock, MagicMock, patch, AsyncMock
from uuid import UUID
from plugins.operators.store_chunks import StoreChunksOperator


class TestStoreChunksOperator:
    """Test suite for StoreChunksOperator."""

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

    @pytest.fixture
    def sample_validate_result(self):
        """Sample validate result data."""
        return {
            'bucket': 'documents',
            'object_key': 'test/sample.xlsx',
            'content_type': 'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet',
            'document_id': '550e8400-e29b-41d4-a716-446655440000',
        }

    @pytest.fixture
    def sample_embedding_result_with_table_chunks(self):
        """Sample embedding result with table chunks."""
        return {
            'embeddings': [
                {'chunk_id': 'doc_t0_r1', 'vector': [0.1] * 3072},
                {'chunk_id': 'doc_t0_r2', 'vector': [0.2] * 3072},
            ],
            'chunks': [
                {
                    'content': {'Name': 'Alice', 'Age': '25'},
                    '_source': {
                        'document_id': '550e8400-e29b-41d4-a716-446655440000',
                        'filename': 'sample.xlsx',
                        'sheet': 'Sheet1',
                        'table_id': 0,
                        'row': 1,
                        'col_range': 'A:B',
                    },
                },
                {
                    'content': {'Name': 'Bob', 'Age': '30'},
                    '_source': {
                        'document_id': '550e8400-e29b-41d4-a716-446655440000',
                        'filename': 'sample.xlsx',
                        'sheet': 'Sheet1',
                        'table_id': 0,
                        'row': 2,
                        'col_range': 'A:B',
                    },
                },
            ],
            'embedding_count': 2,
            'model': 'text-embedding-3-large',
        }

    @pytest.fixture
    def sample_embedding_result_with_semantic_chunks(self):
        """Sample embedding result with semantic chunks."""
        return {
            'embeddings': [
                {'chunk_id': 'semantic_0', 'vector': [0.1] * 3072},
            ],
            'chunks': [
                {
                    'text': 'This is the document content.',
                    'metadata': {'page_number': 1},
                    'original_filename': 'document.docx',
                },
            ],
            'embedding_count': 1,
            'model': 'text-embedding-3-large',
        }

    def test_store_chunks_success(
        self,
        mock_context,
        sample_validate_result,
        sample_embedding_result_with_table_chunks
    ):
        """Test successful chunk storage."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'generate_embeddings':
                return sample_embedding_result_with_table_chunks
            elif task_ids == 'validate_event':
                return sample_validate_result
            return None

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

        with patch('plugins.operators.store_chunks.StructuredStore') as mock_store_class, \
             patch('plugins.operators.store_chunks.async_session_maker') as mock_session_maker, \
             patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:

            # Mock async session
            mock_session = MagicMock()
            mock_session.commit = AsyncMock()
            mock_session.rollback = AsyncMock()
            mock_session.__aenter__ = AsyncMock(return_value=mock_session)
            mock_session.__aexit__ = AsyncMock(return_value=None)
            mock_session_maker.return_value = mock_session

            # Mock store
            mock_store = MagicMock()
            mock_store.store_records = AsyncMock(return_value=2)
            mock_store_class.return_value = mock_store

            # Mock ExtractedRecord
            mock_record_class.side_effect = lambda **kwargs: MagicMock(**kwargs)

            op = StoreChunksOperator(task_id='test_store')

            result = op.execute(mock_context)

            # Verify result structure
            assert result['status'] == 'success'
            assert result['stored_count'] == 2
            assert result['document_id'] == '550e8400-e29b-41d4-a716-446655440000'
            assert result['errors'] == []

    def test_store_chunks_no_chunks(self, mock_context, sample_validate_result):
        """Test handling of empty chunks."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'generate_embeddings':
                return {'chunks': [], 'embeddings': []}
            elif task_ids == 'validate_event':
                return sample_validate_result
            return None

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

        op = StoreChunksOperator(task_id='test_store')

        result = op.execute(mock_context)

        assert result['status'] == 'empty'
        assert result['stored_count'] == 0

    def test_store_chunks_no_embedding_result(self, mock_context, sample_validate_result):
        """Test handling of missing embedding result."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'validate_event':
                return sample_validate_result
            return None

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

        op = StoreChunksOperator(task_id='test_store')

        result = op.execute(mock_context)

        assert result['status'] == 'empty'
        assert result['stored_count'] == 0

    def test_store_chunks_missing_validate_result_raises(self, mock_context):
        """Test that missing validate result raises ValueError."""
        mock_context['ti'].xcom_pull.return_value = None

        op = StoreChunksOperator(task_id='test_store')

        with pytest.raises(ValueError, match='No validate result'):
            op.execute(mock_context)

    def test_store_chunks_semantic_chunks(
        self,
        mock_context,
        sample_validate_result,
        sample_embedding_result_with_semantic_chunks
    ):
        """Test storing semantic chunks."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'generate_embeddings':
                return sample_embedding_result_with_semantic_chunks
            elif task_ids == 'validate_event':
                return sample_validate_result
            return None

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

        with patch('plugins.operators.store_chunks.StructuredStore') as mock_store_class, \
             patch('plugins.operators.store_chunks.async_session_maker') as mock_session_maker, \
             patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:

            # Mock async session
            mock_session = MagicMock()
            mock_session.commit = AsyncMock()
            mock_session.rollback = AsyncMock()
            mock_session.__aenter__ = AsyncMock(return_value=mock_session)
            mock_session.__aexit__ = AsyncMock(return_value=None)
            mock_session_maker.return_value = mock_session

            # Mock store
            mock_store = MagicMock()
            mock_store.store_records = AsyncMock(return_value=1)
            mock_store_class.return_value = mock_store

            # Mock ExtractedRecord
            mock_record_class.side_effect = lambda **kwargs: MagicMock(**kwargs)

            op = StoreChunksOperator(task_id='test_store')

            result = op.execute(mock_context)

            assert result['status'] == 'success'
            assert result['stored_count'] == 1

    def test_custom_task_ids(self):
        """Test using custom task IDs."""
        op = StoreChunksOperator(
            task_id='test_store',
            embeddings_task_id='custom_embeddings',
            validate_task_id='custom_validate',
        )

        assert op.embeddings_task_id == 'custom_embeddings'
        assert op.validate_task_id == 'custom_validate'

    def test_default_constants(self):
        """Test default constant values."""
        assert StoreChunksOperator.DEFAULT_EMBEDDINGS_TASK_ID == 'generate_embeddings'
        assert StoreChunksOperator.DEFAULT_VALIDATE_TASK_ID == 'validate_event'

    def test_empty_result(self):
        """Test _empty_result helper method."""
        op = StoreChunksOperator(task_id='test')
        doc_id = UUID('550e8400-e29b-41d4-a716-446655440000')

        result = op._empty_result(doc_id)

        assert result['stored_count'] == 0
        assert result['document_id'] == '550e8400-e29b-41d4-a716-446655440000'
        assert result['status'] == 'empty'
        assert result['errors'] == []

    def test_prepare_records_table_format(self):
        """Test _prepare_records with table format chunks."""
        op = StoreChunksOperator(task_id='test')

        with patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:
            mock_record_class.side_effect = lambda **kwargs: MagicMock(**kwargs)

            chunks = [
                {
                    'content': {'Name': 'Alice'},
                    '_source': {
                        'document_id': '550e8400-e29b-41d4-a716-446655440000',
                        'filename': 'test.xlsx',
                        'sheet': 'Sheet1',
                        'table_id': 0,
                        'row': 1,
                        'col_range': 'A:A',
                    },
                }
            ]

            doc_id = UUID('550e8400-e29b-41d4-a716-446655440000')
            records = op._prepare_records(chunks, doc_id)

            assert len(records) == 1
            mock_record_class.assert_called_once()

    def test_prepare_records_semantic_format(self):
        """Test _prepare_records with semantic format chunks."""
        op = StoreChunksOperator(task_id='test')

        with patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:
            mock_record_class.side_effect = lambda **kwargs: MagicMock(**kwargs)

            chunks = [
                {
                    'text': 'Document content',
                    'metadata': {'page': 1},
                    'original_filename': 'doc.docx',
                }
            ]

            doc_id = UUID('550e8400-e29b-41d4-a716-446655440000')
            records = op._prepare_records(chunks, doc_id)

            assert len(records) == 1

    def test_prepare_records_unknown_format(self):
        """Test _prepare_records skips unknown format."""
        op = StoreChunksOperator(task_id='test')

        with patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:
            chunks = [
                {'unknown_field': 'value'},  # Unknown format
            ]

            doc_id = UUID('550e8400-e29b-41d4-a716-446655440000')
            records = op._prepare_records(chunks, doc_id)

            # Unknown format should be skipped
            assert len(records) == 0

    def test_storage_error_propagation(
        self,
        mock_context,
        sample_validate_result,
        sample_embedding_result_with_table_chunks
    ):
        """Test that storage errors are properly propagated."""
        def xcom_pull_side_effect(task_ids):
            if task_ids == 'generate_embeddings':
                return sample_embedding_result_with_table_chunks
            elif task_ids == 'validate_event':
                return sample_validate_result
            return None

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

        with patch('plugins.operators.store_chunks.StructuredStore') as mock_store_class, \
             patch('plugins.operators.store_chunks.async_session_maker') as mock_session_maker, \
             patch('plugins.operators.store_chunks.ExtractedRecord') as mock_record_class:

            # Mock async session (rollback handled by context manager)
            mock_session = MagicMock()
            mock_session.__aenter__ = AsyncMock(return_value=mock_session)
            mock_session.__aexit__ = AsyncMock(return_value=None)
            mock_session_maker.return_value = mock_session

            # Mock store to raise error
            mock_store = MagicMock()
            mock_store.store_records = AsyncMock(side_effect=Exception('Database error'))
            mock_store_class.return_value = mock_store

            # Mock ExtractedRecord
            mock_record_class.side_effect = lambda **kwargs: MagicMock(**kwargs)

            op = StoreChunksOperator(task_id='test_store')

            with pytest.raises(Exception, match='Database error'):
                op.execute(mock_context)
