"""
Unit tests for DetectHeadersOperator.
"""

import pytest
from unittest.mock import Mock, MagicMock, patch
from plugins.operators.detect_headers import DetectHeadersOperator


class TestDetectHeadersOperator:
    """Test suite for DetectHeadersOperator."""

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

    @pytest.fixture
    def sample_excel_parse_result(self):
        """Sample Excel parse result data."""
        return {
            'doc_type': 'excel',
            'html_content': '<html><table>...</table></html>',
            'markdown_content': None,
            'tables': [
                {
                    'html': '<table><tr><th>Name</th><th>Age</th></tr><tr><td>Alice</td><td>25</td></tr></table>',
                    'rows': [['Name', 'Age'], ['Alice', '25']],
                },
                {
                    'html': '<table><tr><th>Product</th><th>Price</th></tr><tr><td>Widget</td><td>100</td></tr></table>',
                    'rows': [['Product', 'Price'], ['Widget', '100']],
                }
            ],
            'raw_text_tables': [],
            'local_path': '/tmp/test/sample.xlsx',
            'document_id': 'test-doc-123',
        }

    def test_detect_headers_excel_success(self, mock_context, sample_excel_parse_result):
        """Test successful header detection for Excel document."""
        mock_context['ti'].xcom_pull.return_value = sample_excel_parse_result

        with patch('plugins.operators.detect_headers.LlmHeaderDetector') as mock_detector_class:
            # Mock detector
            mock_detector = MagicMock()
            mock_header_result = MagicMock()
            mock_header_result.header_rows = [0]
            mock_header_result.headers = [
                {'column': 0, 'hierarchy': ['Name']},
                {'column': 1, 'hierarchy': ['Age']},
            ]
            mock_detector.detect_headers.return_value = mock_header_result
            mock_detector_class.return_value = mock_detector

            op = DetectHeadersOperator(
                task_id='test_detect',
                parse_task_id='parse_excel',
            )

            result = op.execute(mock_context)

            # Verify XCom pull
            mock_context['ti'].xcom_pull.assert_called_once_with(task_ids='parse_excel')

            # Verify result structure
            assert result['tables_processed'] == 2
            assert len(result['header_data']) == 2

            # Verify first table header data
            first_table = result['header_data'][0]
            assert first_table['table_index'] == 0
            assert first_table['header_rows'] == [0]
            assert len(first_table['headers']) == 2

    def test_detect_headers_non_excel_skips(self, mock_context):
        """Test that non-Excel documents are skipped."""
        mock_context['ti'].xcom_pull.return_value = {
            'doc_type': 'word',
            'html_content': None,
            'markdown_content': '# Document',
            'tables': [],
        }

        op = DetectHeadersOperator(task_id='test_detect')
        result = op.execute(mock_context)

        assert result['tables_processed'] == 0
        assert result['header_data'] == []

    def test_detect_headers_no_parse_result(self, mock_context):
        """Test handling of missing parse result."""
        mock_context['ti'].xcom_pull.return_value = None

        op = DetectHeadersOperator(task_id='test_detect')
        result = op.execute(mock_context)

        assert result['tables_processed'] == 0
        assert result['header_data'] == []

    def test_detect_headers_no_tables(self, mock_context):
        """Test handling of Excel with no tables."""
        mock_context['ti'].xcom_pull.return_value = {
            'doc_type': 'excel',
            'tables': [],
        }

        op = DetectHeadersOperator(task_id='test_detect')
        result = op.execute(mock_context)

        assert result['tables_processed'] == 0
        assert result['header_data'] == []

    def test_detect_headers_with_fallback(self, mock_context, sample_excel_parse_result):
        """Test fallback when LLM detection fails."""
        mock_context['ti'].xcom_pull.return_value = sample_excel_parse_result

        with patch('plugins.operators.detect_headers.LlmHeaderDetector') as mock_detector_class:
            # Mock detector that raises exception
            mock_detector = MagicMock()
            mock_detector.detect_headers.side_effect = Exception('LLM error')
            mock_detector_class.return_value = mock_detector

            op = DetectHeadersOperator(task_id='test_detect')
            result = op.execute(mock_context)

            # Should still return results with fallback headers
            assert result['tables_processed'] == 2
            assert len(result['header_data']) == 2

            # Verify fallback headers (first row as default)
            for header_info in result['header_data']:
                assert header_info['header_rows'] == [0]
                assert header_info['headers'] == []

    def test_detect_headers_empty_html_fallback(self, mock_context):
        """Test fallback when table has empty HTML."""
        mock_context['ti'].xcom_pull.return_value = {
            'doc_type': 'excel',
            'tables': [
                {'html': '', 'rows': [['A', 'B'], ['1', '2']]},
            ],
        }

        op = DetectHeadersOperator(task_id='test_detect')
        result = op.execute(mock_context)

        assert result['tables_processed'] == 1
        # Should use fallback
        assert result['header_data'][0]['header_rows'] == [0]
        assert result['header_data'][0]['headers'] == []

    def test_detect_headers_multi_level(self, mock_context):
        """Test detection of multi-level headers."""
        mock_context['ti'].xcom_pull.return_value = {
            'doc_type': 'excel',
            'tables': [{
                'html': '<table><tr><th colspan="2">Sales</th></tr><tr><th>Q1</th><th>Q2</th></tr></table>',
                'rows': [['Sales', ''], ['Q1', 'Q2'], ['100', '200']],
            }],
        }

        with patch('plugins.operators.detect_headers.LlmHeaderDetector') as mock_detector_class:
            mock_detector = MagicMock()
            mock_header_result = MagicMock()
            mock_header_result.header_rows = [0, 1]
            mock_header_result.headers = [
                {'column': 0, 'hierarchy': ['Sales', 'Q1']},
                {'column': 1, 'hierarchy': ['Sales', 'Q2']},
            ]
            mock_detector.detect_headers.return_value = mock_header_result
            mock_detector_class.return_value = mock_detector

            op = DetectHeadersOperator(task_id='test_detect')
            result = op.execute(mock_context)

            assert result['tables_processed'] == 1
            assert result['header_data'][0]['header_rows'] == [0, 1]
            assert len(result['header_data'][0]['headers']) == 2

    def test_custom_parse_task_id(self, mock_context, sample_excel_parse_result):
        """Test using custom parse task ID."""
        mock_context['ti'].xcom_pull.return_value = sample_excel_parse_result

        with patch('plugins.operators.detect_headers.LlmHeaderDetector') as mock_detector_class:
            mock_detector = MagicMock()
            mock_header_result = MagicMock()
            mock_header_result.header_rows = [0]
            mock_header_result.headers = []
            mock_detector.detect_headers.return_value = mock_header_result
            mock_detector_class.return_value = mock_detector

            op = DetectHeadersOperator(
                task_id='test_detect',
                parse_task_id='custom_parse'
            )

            op.execute(mock_context)

            # Verify custom task_id was used
            mock_context['ti'].xcom_pull.assert_called_once_with(task_ids='custom_parse')

    def test_operator_default_parse_task_id(self):
        """Test default parse task ID constant."""
        op = DetectHeadersOperator(task_id='test')
        assert op.parse_task_id == 'parse_excel'
