"""Unit tests for multi-pass pipeline orchestration (Story 3.8a).

Tests the 4-pass extraction architecture:
    PASS 1: Symbol Detection
    PASS 2: Table Extraction
    PASS 3: Symbol Resolution
    PASS 4: Cross-Reference Detection
"""

from unittest.mock import MagicMock, Mock, patch
from uuid import uuid4

import pytest

from src.extraction.cross_ref_detector import CrossRefResolutionSummary
from src.extraction.models import (
    ExtractedRecord,
    SymbolDetectionSummary,
    SymbolResolutionSummary,
)
from src.extraction.pipeline import ExtractionPipeline, ExtractionSummary


class TestMultiPassPipelineOrchestration:
    """Test multi-pass pipeline execution order and coordination."""

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

    @pytest.fixture
    def mock_document(self):
        """Create a mock document."""
        from src.db import models as db_models

        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
        doc.api_key_id = uuid4()
        return doc

    @pytest.fixture
    def empty_summaries(self, mock_document):
        """Create empty summaries for non-critical passes."""
        symbol_detection = SymbolDetectionSummary(
            dictionaries_found=0,
            total_symbols=0,
            detected_sheets=[],
            warnings=[],
        )
        symbol_resolution = SymbolResolutionSummary(
            document_id=mock_document.id,
            records_processed=0,
            symbols_resolved=0,
            unresolved_count=0,
            unresolved_symbols=[],
            warnings=[],
        )
        cross_ref = CrossRefResolutionSummary(
            document_id=mock_document.id,
            references_found=0,
            references_resolved=0,
            unresolved_count=0,
            unresolved_refs=[],
            warnings=[],
        )
        return symbol_detection, symbol_resolution, cross_ref

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_pass_execution_order(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document, empty_summaries
    ):
        """Test that passes execute in correct order: 1 → 2 → 3 → 4."""
        # 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

        symbol_detection, symbol_resolution, cross_ref = empty_summaries

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Track execution order
        execution_order = []

        def track_symbol_detection(*args, **kwargs):
            execution_order.append("PASS1_SYMBOL_DETECTION")
            return ([], symbol_detection)

        def track_table_detection(*args, **kwargs):
            execution_order.append("PASS2_TABLE_EXTRACTION")
            return []

        def track_symbol_resolution(*args, **kwargs):
            execution_order.append("PASS3_SYMBOL_RESOLUTION")
            return symbol_resolution

        def track_cross_ref_detection(*args, **kwargs):
            execution_order.append("PASS4_CROSS_REF_DETECTION")
            return cross_ref

        with patch.object(
            pipeline.symbol_detector, "detect_dictionaries", side_effect=track_symbol_detection
        ), patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", side_effect=track_table_detection
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", side_effect=track_symbol_resolution
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", side_effect=track_cross_ref_detection
        ):
            pipeline.process_document(mock_document.id)

        # Verify execution order
        assert execution_order == [
            "PASS1_SYMBOL_DETECTION",
            "PASS2_TABLE_EXTRACTION",
            "PASS3_SYMBOL_RESOLUTION",
            "PASS4_CROSS_REF_DETECTION",
        ]

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_pass1_symbol_detection_runs_before_extraction(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document, empty_summaries
    ):
        """Test that PASS 1 symbol detection runs before PASS 2 table extraction."""
        # 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

        symbol_detection, symbol_resolution, cross_ref = empty_summaries

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Track if symbol detection ran before table detection
        symbol_detection_completed = [False]

        def mark_symbol_detection_done(*args, **kwargs):
            symbol_detection_completed[0] = True
            return ([], symbol_detection)

        def check_symbol_detection_ran(*args, **kwargs):
            assert symbol_detection_completed[0], "Symbol detection should run before table extraction"
            return []

        with patch.object(
            pipeline.symbol_detector, "detect_dictionaries", side_effect=mark_symbol_detection_done
        ), patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", side_effect=check_symbol_detection_ran
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", return_value=symbol_resolution
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", return_value=cross_ref
        ):
            pipeline.process_document(mock_document.id)

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_pass1_failure_doesnt_fail_pipeline(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document, empty_summaries
    ):
        """Test that PASS 1 symbol detection failure doesn't fail the entire pipeline."""
        # 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

        _, symbol_resolution, cross_ref = empty_summaries

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Make symbol detection fail
        with patch.object(
            pipeline.symbol_detector, "detect_dictionaries", side_effect=Exception("Symbol detection failed")
        ), patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", return_value=[]
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", return_value=symbol_resolution
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", return_value=cross_ref
        ):
            # Should not raise - PASS 1 is non-critical
            result = pipeline.process_document(mock_document.id)

        # Pipeline should complete with warning
        assert isinstance(result, ExtractionSummary)
        assert mock_document.status == "completed"
        assert any("Symbol detection failed" in w for w in result.warnings)

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_pass3_failure_doesnt_fail_pipeline(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document, empty_summaries
    ):
        """Test that PASS 3 symbol resolution failure doesn't fail the entire pipeline."""
        # 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

        symbol_detection, _, cross_ref = empty_summaries

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Make symbol resolution fail
        with patch.object(
            pipeline.symbol_detector, "detect_dictionaries", return_value=([], symbol_detection)
        ), patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", return_value=[]
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", side_effect=Exception("Resolution failed")
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", return_value=cross_ref
        ):
            # Should not raise - PASS 3 is non-critical
            result = pipeline.process_document(mock_document.id)

        # Pipeline should complete with warning
        assert isinstance(result, ExtractionSummary)
        assert mock_document.status == "completed"
        assert any("Symbol resolution failed" in w for w in result.warnings)

    @patch("src.extraction.pipeline.openpyxl.load_workbook")
    @patch("src.extraction.pipeline.Path")
    def test_pass4_failure_doesnt_fail_pipeline(
        self, mock_path, mock_load_workbook, mock_db_session, mock_document, empty_summaries
    ):
        """Test that PASS 4 cross-reference detection failure doesn't fail the entire pipeline."""
        # 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

        symbol_detection, symbol_resolution, _ = empty_summaries

        pipeline = ExtractionPipeline(db_session=mock_db_session)

        # Make cross-ref detection fail
        with patch.object(
            pipeline.symbol_detector, "detect_dictionaries", return_value=([], symbol_detection)
        ), patch.object(
            pipeline.table_detector, "_detect_boundaries_in_sheet", return_value=[]
        ), patch.object(
            pipeline.symbol_resolver, "resolve_document", return_value=symbol_resolution
        ), patch.object(
            pipeline.cross_ref_detector, "process_document", side_effect=Exception("Cross-ref failed")
        ):
            # Should not raise - PASS 4 is non-critical
            result = pipeline.process_document(mock_document.id)

        # Pipeline should complete with warning
        assert isinstance(result, ExtractionSummary)
        assert mock_document.status == "completed"
        assert any("Cross-reference detection failed" in w for w in result.warnings)


class TestExtractionSummaryStatistics:
    """Test that ExtractionSummary includes all pass statistics."""

    def test_summary_includes_symbol_detection_stats(self):
        """Test that ExtractionSummary includes symbol detection stats."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=5000,
            warnings=[],
            symbols_detected=15,
        )

        assert summary.symbols_detected == 15
        d = summary.to_dict()
        assert d["symbols_detected"] == 15

    def test_summary_includes_symbol_resolution_stats(self):
        """Test that ExtractionSummary includes symbol resolution stats."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=5000,
            warnings=[],
            symbols_resolved=142,
            unresolved_symbols=3,
        )

        assert summary.symbols_resolved == 142
        assert summary.unresolved_symbols == 3
        d = summary.to_dict()
        assert d["symbols_resolved"] == 142
        assert d["unresolved_symbols"] == 3

    def test_summary_includes_cross_ref_stats(self):
        """Test that ExtractionSummary includes cross-reference stats."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=5000,
            warnings=[],
            cross_references_found=45,
            cross_references_resolved=38,
            cross_references_unresolved=7,
        )

        assert summary.cross_references_found == 45
        assert summary.cross_references_resolved == 38
        assert summary.cross_references_unresolved == 7
        d = summary.to_dict()
        assert d["cross_references_found"] == 45
        assert d["cross_references_resolved"] == 38
        assert d["cross_references_unresolved"] == 7

    def test_summary_defaults_all_to_zero(self):
        """Test that all new stats default to 0."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=1,
            tables_detected=1,
            records_extracted=10,
            duration_ms=1000,
            warnings=[],
        )

        assert summary.symbols_detected == 0
        assert summary.symbols_resolved == 0
        assert summary.unresolved_symbols == 0
        assert summary.cross_references_found == 0
        assert summary.cross_references_resolved == 0
        assert summary.cross_references_unresolved == 0


class TestProgressTracking:
    """Test progress tracking for multi-pass pipeline."""

    @pytest.fixture
    def mock_redis_client(self):
        """Create a mock Redis client."""
        client = Mock()
        client.setex = Mock()
        return client

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

    def test_update_stage_progress(self, mock_db_session, mock_redis_client):
        """Test _update_stage_progress updates Redis with stage info."""
        import json

        pipeline = ExtractionPipeline(db_session=mock_db_session, redis_client=mock_redis_client)
        doc_id = uuid4()

        pipeline._update_stage_progress(doc_id, "detecting_symbols")

        mock_redis_client.setex.assert_called_once()
        call_args = mock_redis_client.setex.call_args
        key = call_args[0][0]
        ttl = call_args[0][1]
        data = json.loads(call_args[0][2])

        assert key == f"processing:documents:{doc_id}"
        assert ttl == 3600
        assert data["status"] == "processing"
        assert data["stage"] == "detecting_symbols"

    def test_update_stage_progress_with_extra_data(self, mock_db_session, mock_redis_client):
        """Test _update_stage_progress includes extra data."""
        import json

        pipeline = ExtractionPipeline(db_session=mock_db_session, redis_client=mock_redis_client)
        doc_id = uuid4()

        pipeline._update_stage_progress(
            doc_id,
            "extracting_tables",
            extra_data={"sheets_processed": 2, "records_count": 100},
        )

        call_args = mock_redis_client.setex.call_args
        data = json.loads(call_args[0][2])

        assert data["stage"] == "extracting_tables"
        assert data["sheets_processed"] == 2
        assert data["records_count"] == 100

    def test_update_stage_progress_no_redis(self, mock_db_session):
        """Test _update_stage_progress does nothing without Redis client."""
        pipeline = ExtractionPipeline(db_session=mock_db_session, redis_client=None)
        doc_id = uuid4()

        # Should not raise
        pipeline._update_stage_progress(doc_id, "detecting_symbols")

    def test_update_stage_progress_redis_error_handled(self, mock_db_session, mock_redis_client):
        """Test _update_stage_progress handles Redis errors gracefully."""
        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 - progress tracking is non-critical
        pipeline._update_stage_progress(doc_id, "detecting_symbols")
