"""Unit tests for extraction pipeline orchestrator."""

import json
from datetime import datetime, timezone
from pathlib import Path
from unittest.mock import MagicMock, Mock, patch
from uuid import UUID, uuid4

import pytest
from openpyxl import Workbook
from openpyxl.utils.exceptions import InvalidFileException

from src.core.exceptions import ExtractionError
from src.db import models as db_models
from src.extraction.models import (
    ExtractedRecord,
    HeaderInfo,
    MergeMetadata,
    TableBoundary,
)
from src.extraction.pipeline import ExtractionPipeline, ExtractionSummary


@pytest.fixture
def mock_db_session():
    """Mock database session."""
    session = Mock()
    session.commit = Mock()
    session.rollback = Mock()
    session.scalar = Mock()
    session.bulk_insert_mappings = Mock()
    return session


@pytest.fixture
def mock_redis_client():
    """Mock Redis client."""
    redis_client = Mock()
    redis_client.setex = Mock()
    return redis_client


@pytest.fixture
def mock_document():
    """Mock document model."""
    doc = Mock(spec=db_models.Document)
    doc.id = uuid4()
    doc.filename = "test_document.xlsx"
    doc.file_path = "/tmp/test_document.xlsx"
    doc.status = "pending"
    doc.sheet_count = None
    doc.record_count = None
    doc.processed_at = None
    doc.error_message = None
    return doc


@pytest.fixture
def sample_extraction_record():
    """Create a sample extraction record."""
    return ExtractedRecord(
        content={"Column A": "Value 1", "Column B": "Value 2"},
        headers=["Column A", "Column B"],
        _source={
            "document_id": str(uuid4()),
            "filename": "test.xlsx",
            "sheet": "Sheet1",
            "table_id": 1,
            "row": 2,
            "col_range": "A:B",
        },
    )


class TestExtractionSummary:
    """Tests for ExtractionSummary dataclass."""

    def test_to_dict(self):
        """Test ExtractionSummary.to_dict() serialization."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=1500,
            warnings=["Warning 1", "Warning 2"],
        )

        result = summary.to_dict()

        assert result["document_id"] == str(doc_id)
        assert result["sheets_processed"] == 3
        assert result["tables_detected"] == 5
        assert result["records_extracted"] == 100
        assert result["duration_ms"] == 1500
        assert result["warnings"] == ["Warning 1", "Warning 2"]


class TestExtractionPipelineInit:
    """Tests for ExtractionPipeline initialization."""

    def test_init_with_required_params(self, mock_db_session):
        """Test pipeline initialization with required parameters."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        assert pipeline.db_session == mock_db_session
        assert pipeline.redis_client is None
        assert pipeline.table_detector is not None
        assert pipeline.header_extractor is not None
        assert pipeline.merged_cell_resolver is not None
        assert pipeline.record_builder is not None

    def test_init_with_redis(self, mock_db_session, mock_redis_client):
        """Test pipeline initialization with Redis client."""
        pipeline = ExtractionPipeline(
            db_session=mock_db_session, redis_client=mock_redis_client
        )

        assert pipeline.redis_client == mock_redis_client


class TestLoadDocument:
    """Tests for _load_document method."""

    def test_load_document_success(self, mock_db_session, mock_document):
        """Test successful document loading."""
        mock_db_session.scalar.return_value = mock_document
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        result = pipeline._load_document(mock_document.id)

        assert result == mock_document
        mock_db_session.scalar.assert_called_once()

    def test_load_document_not_found(self, mock_db_session):
        """Test document not found error."""
        mock_db_session.scalar.return_value = None
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        doc_id = uuid4()

        with pytest.raises(ExtractionError) as exc_info:
            pipeline._load_document(doc_id)

        assert f"Document not found: {doc_id}" in str(exc_info.value)


class TestLoadWorkbook:
    """Tests for _load_workbook method."""

    def test_load_workbook_file_not_found(self, mock_db_session):
        """Test workbook loading with non-existent file."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        with pytest.raises(ExtractionError) as exc_info:
            pipeline._load_workbook("/nonexistent/file.xlsx")

        assert "File not found" in str(exc_info.value)

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_load_workbook_invalid_format(
        self, mock_path, mock_load_workbook, mock_db_session
    ):
        """Test workbook loading with invalid file format."""
        # Mock Path.exists to return True
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_load_workbook.side_effect = InvalidFileException("Invalid file")
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        with pytest.raises(ExtractionError) as exc_info:
            pipeline._load_workbook("/tmp/invalid.xlsx")

        assert "Invalid Excel file format" in str(exc_info.value)

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_load_workbook_success(
        self, mock_path, mock_load_workbook, mock_db_session
    ):
        """Test successful workbook loading."""
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_workbook = Mock()
        mock_load_workbook.return_value = mock_workbook

        pipeline = ExtractionPipeline(db_session=mock_db_session)
        result = pipeline._load_workbook("/tmp/test.xlsx")

        assert result == mock_workbook
        mock_load_workbook.assert_called_once_with("/tmp/test.xlsx", data_only=True)


class TestUpdateDocumentStatus:
    """Tests for _update_document_status method."""

    def test_update_status_to_processing(self, mock_db_session, mock_document):
        """Test updating document status to processing."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        pipeline._update_document_status(mock_document, "processing")

        assert mock_document.status == "processing"
        mock_db_session.commit.assert_called_once()

    def test_update_status_to_completed(self, mock_db_session, mock_document):
        """Test updating document status to completed."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)

        pipeline._update_document_status(
            mock_document, "completed", sheet_count=3, record_count=100
        )

        assert mock_document.status == "completed"
        assert mock_document.sheet_count == 3
        assert mock_document.record_count == 100
        assert isinstance(mock_document.processed_at, datetime)
        mock_db_session.commit.assert_called_once()

    def test_update_status_to_failed(self, mock_db_session, mock_document):
        """Test updating document status to failed."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        error_msg = "Extraction failed due to invalid format"

        pipeline._update_document_status(
            mock_document, "failed", error_message=error_msg
        )

        assert mock_document.status == "failed"
        assert mock_document.error_message == error_msg
        mock_db_session.commit.assert_called_once()


class TestStoreRecords:
    """Tests for _store_records method."""

    def test_store_records_empty_list(self, mock_db_session):
        """Test storing empty list of records."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        doc_id = uuid4()

        pipeline._store_records([], doc_id)

        mock_db_session.bulk_insert_mappings.assert_not_called()

    def test_store_records_single_batch(
        self, mock_db_session, sample_extraction_record
    ):
        """Test storing records in single batch."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        doc_id = uuid4()
        records = [sample_extraction_record]

        pipeline._store_records(records, doc_id)

        mock_db_session.bulk_insert_mappings.assert_called_once()
        mock_db_session.commit.assert_called_once()

    def test_store_records_multiple_batches(
        self, mock_db_session, sample_extraction_record
    ):
        """Test storing records in multiple batches."""
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        doc_id = uuid4()

        # Create 1000 records to trigger multiple batches (batch_size=500)
        records = [sample_extraction_record] * 1000

        pipeline._store_records(records, doc_id)

        # Should be called twice (1000 / 500 = 2)
        assert mock_db_session.bulk_insert_mappings.call_count == 2
        assert mock_db_session.commit.call_count == 2

    def test_store_records_database_error(
        self, mock_db_session, sample_extraction_record
    ):
        """Test storing records with database error."""
        mock_db_session.bulk_insert_mappings.side_effect = Exception(
            "Database error"
        )
        pipeline = ExtractionPipeline(db_session=mock_db_session)
        doc_id = uuid4()

        with pytest.raises(ExtractionError) as exc_info:
            pipeline._store_records([sample_extraction_record], doc_id)

        assert "Failed to store extracted records" in str(exc_info.value)
        mock_db_session.rollback.assert_called_once()


class TestUpdateProgress:
    """Tests for _update_progress method."""

    def test_update_progress_no_redis(self, mock_db_session):
        """Test progress update without Redis client."""
        pipeline = ExtractionPipeline(db_session=mock_db_session, redis_client=None)
        doc_id = uuid4()

        # Should not raise exception
        pipeline._update_progress(doc_id, 1, 3, 50)

    def test_update_progress_with_redis(self, mock_db_session, mock_redis_client):
        """Test progress update with Redis client."""
        pipeline = ExtractionPipeline(
            db_session=mock_db_session, redis_client=mock_redis_client
        )
        doc_id = uuid4()

        pipeline._update_progress(doc_id, 2, 3, 75)

        mock_redis_client.setex.assert_called_once()
        args = mock_redis_client.setex.call_args
        assert args[0][0] == f"processing:documents:{doc_id}"
        assert args[0][1] == 3600  # TTL

        progress_data = json.loads(args[0][2])
        assert progress_data["status"] == "processing"
        assert progress_data["current_sheet"] == 2
        assert progress_data["total_sheets"] == 3
        assert progress_data["records"] == 75

    def test_update_progress_redis_error(self, mock_db_session, mock_redis_client):
        """Test progress update with Redis error (should not fail)."""
        mock_redis_client.setex.side_effect = Exception("Redis connection error")
        pipeline = ExtractionPipeline(
            db_session=mock_db_session, redis_client=mock_redis_client
        )
        doc_id = uuid4()

        # Should not raise exception - graceful degradation
        pipeline._update_progress(doc_id, 1, 3, 50)


class TestProcessDocument:
    """Tests for process_document method."""

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_process_document_empty_workbook(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document
    ):
        """Test processing document with empty workbook."""
        # Setup mocks
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_workbook = Mock()
        mock_workbook.worksheets = []
        mock_load_workbook.return_value = mock_workbook

        mock_db_session.scalar.return_value = mock_document

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Execute
        result = pipeline.process_document(mock_document.id)

        # Verify
        assert isinstance(result, ExtractionSummary)
        assert result.sheets_processed == 0
        assert result.tables_detected == 0
        assert result.records_extracted == 0
        assert len(result.warnings) == 1
        assert "no sheets" in result.warnings[0].lower()
        assert mock_document.status == "completed"
        assert mock_document.sheet_count == 0
        assert mock_document.record_count == 0

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_process_document_no_tables(
        self,
        mock_path,
        mock_load_workbook,
        mock_db_session,
        mock_document,
    ):
        """Test processing document with sheets but no tables."""
        from src.extraction.models import SymbolDetectionSummary, SymbolResolutionSummary
        from src.extraction.cross_ref_detector import CrossRefResolutionSummary

        # Setup mocks
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_sheet = Mock()
        mock_sheet.title = "Sheet1"

        mock_workbook = Mock()
        mock_workbook.worksheets = [mock_sheet]
        mock_load_workbook.return_value = mock_workbook

        mock_db_session.scalar.return_value = mock_document

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Mock all extraction components to return empty results
        empty_symbol_summary = SymbolDetectionSummary(
            dictionaries_found=0, total_symbols=0, detected_sheets=[], warnings=[]
        )
        empty_resolution_summary = SymbolResolutionSummary(
            document_id=mock_document.id,
            records_processed=0,
            symbols_resolved=0,
            unresolved_count=0,
            unresolved_symbols=[],
            warnings=[],
        )
        empty_crossref_summary = CrossRefResolutionSummary(
            document_id=mock_document.id,
            references_found=0,
            references_resolved=0,
            unresolved_count=0,
            unresolved_refs=[],
            warnings=[],
        )

        with patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", return_value=[]
        ), patch.object(
            pipeline.symbol_detector, "detect_dictionaries", return_value=([], empty_symbol_summary)
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", return_value=empty_resolution_summary
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", return_value=empty_crossref_summary
        ):
            result = pipeline.process_document(mock_document.id)

        # Verify
        assert isinstance(result, ExtractionSummary)
        assert result.sheets_processed == 1
        assert result.tables_detected == 0
        assert result.records_extracted == 0
        # Only expect one warning about no tables (non-critical passes are mocked)
        table_warnings = [w for w in result.warnings if "No tables detected" in w]
        assert len(table_warnings) == 1

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_process_document_success(
        self,
        mock_path,
        mock_load_workbook,
        mock_db_session,
        mock_document,
        sample_extraction_record,
    ):
        """Test successful document processing."""
        from src.extraction.models import SymbolDetectionSummary, SymbolResolutionSummary
        from src.extraction.cross_ref_detector import CrossRefResolutionSummary

        # Setup mocks
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_sheet = Mock()
        mock_sheet.title = "Sheet1"

        mock_workbook = Mock()
        mock_workbook.worksheets = [mock_sheet]
        mock_load_workbook.return_value = mock_workbook

        mock_db_session.scalar.return_value = mock_document

        # Create test boundary dict (internal format)
        boundary_dict = {
            "start_row": 1,
            "end_row": 3,
            "start_col": "A",
            "end_col": "B",
        }

        # Create test headers
        headers = {
            "A": HeaderInfo(display="Column A", levels=["Column A"], depth=1, column_letter="A"),
            "B": HeaderInfo(display="Column B", levels=["Column B"], depth=1, column_letter="B"),
        }

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Mock multi-pass components to return empty summaries
        empty_symbol_summary = SymbolDetectionSummary(
            dictionaries_found=0, total_symbols=0, detected_sheets=[], warnings=[]
        )
        empty_resolution_summary = SymbolResolutionSummary(
            document_id=mock_document.id,
            records_processed=0,
            symbols_resolved=0,
            unresolved_count=0,
            unresolved_symbols=[],
            warnings=[],
        )
        empty_crossref_summary = CrossRefResolutionSummary(
            document_id=mock_document.id,
            references_found=0,
            references_resolved=0,
            unresolved_count=0,
            unresolved_refs=[],
            warnings=[],
        )

        # Mock all extraction components
        with patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", return_value=[boundary_dict]
        ), patch.object(
            pipeline.header_extractor, "extract_headers", return_value=headers
        ), patch.object(
            pipeline.merged_cell_resolver,
            "resolve_merged_cells",
            return_value=([[]], []),
        ), patch.object(
            pipeline.record_builder, "build_records", return_value=[sample_extraction_record]
        ), patch.object(
            pipeline.symbol_detector, "detect_dictionaries", return_value=([], empty_symbol_summary)
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", return_value=empty_resolution_summary
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", return_value=empty_crossref_summary
        ):
            result = pipeline.process_document(mock_document.id)

        # Verify
        assert isinstance(result, ExtractionSummary)
        assert result.sheets_processed == 1
        assert result.tables_detected == 1
        assert result.records_extracted == 1
        assert result.warnings == []
        assert mock_document.status == "completed"
        assert mock_document.record_count == 1

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_process_document_extraction_failure(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document
    ):
        """Test document processing with extraction failure."""
        # Setup mocks
        mock_path_instance = MagicMock()
        mock_path_instance.exists.return_value = True
        mock_path.return_value = mock_path_instance

        mock_workbook = Mock()
        mock_load_workbook.side_effect = Exception("Unexpected error")

        mock_db_session.scalar.return_value = mock_document

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Execute - should raise ExtractionError
        with pytest.raises(ExtractionError) as exc_info:
            pipeline.process_document(mock_document.id)

        assert "Extraction pipeline failed" in str(exc_info.value)
        assert mock_document.status == "failed"
        assert mock_document.error_message is not None
