"""Integration tests for multi-pass extraction pipeline (Story 3.8a).

Tests the complete document processing flow with all four passes:
    PASS 1: Symbol Detection
    PASS 2: Table Extraction
    PASS 3: Symbol Resolution
    PASS 4: Cross-Reference Detection
"""

from unittest.mock import Mock
from uuid import uuid4

import pytest

from src.extraction.pipeline import ExtractionPipeline, ExtractionSummary


class TestMultiPassPipelineIntegration:
    """Integration tests for multi-pass pipeline with real Excel files."""

    @pytest.fixture
    def fixtures_dir(self):
        """Get path to test fixtures directory."""
        from pathlib import Path

        return Path(__file__).parent.parent / "fixtures" / "sample_excel"

    @pytest.fixture
    def mock_db_session(self):
        """Create a mock database session with proper query chain."""
        session = Mock()
        session.commit = Mock()
        session.rollback = Mock()
        session.bulk_insert_mappings = Mock()

        # Mock document
        from src.db import models as db_models

        doc = Mock(spec=db_models.Document)
        doc.id = uuid4()
        doc.filename = "test_standard_table.xlsx"
        doc.file_path = None  # Will be set in test
        doc.status = "pending"
        doc.sheet_count = None
        doc.record_count = None
        doc.processed_at = None
        doc.error_message = None
        doc.api_key_id = uuid4()

        session.scalar = Mock(return_value=doc)

        # Mock query chain for symbol resolution and cross-ref detection
        mock_query = Mock()
        session.query = Mock(return_value=mock_query)
        mock_query.filter = Mock(return_value=mock_query)
        mock_query.order_by = Mock(return_value=mock_query)
        mock_query.all = Mock(return_value=[])
        mock_query.first = Mock(return_value=None)

        return session, doc

    def test_full_pipeline_with_all_passes(self, fixtures_dir, mock_db_session):
        """Test complete pipeline execution with all four passes."""
        session, doc = mock_db_session
        file_path = fixtures_dir / "test_standard_table.xlsx"
        doc.file_path = str(file_path)

        pipeline = ExtractionPipeline(db_session=session)
        summary = pipeline.process_document(doc.id)

        # Verify summary is complete
        assert isinstance(summary, ExtractionSummary)
        assert summary.document_id == doc.id
        assert summary.sheets_processed >= 1
        assert summary.duration_ms > 0

        # Verify all statistics are included in summary
        assert hasattr(summary, "symbols_detected")
        assert hasattr(summary, "symbols_resolved")
        assert hasattr(summary, "unresolved_symbols")
        assert hasattr(summary, "cross_references_found")
        assert hasattr(summary, "cross_references_resolved")
        assert hasattr(summary, "cross_references_unresolved")

        # Verify to_dict includes all fields
        d = summary.to_dict()
        assert "symbols_detected" in d
        assert "symbols_resolved" in d
        assert "unresolved_symbols" in d
        assert "cross_references_found" in d
        assert "cross_references_resolved" in d
        assert "cross_references_unresolved" in d

    def test_pipeline_with_merged_cells(self, fixtures_dir, mock_db_session):
        """Test pipeline handles merged cells correctly through all passes."""
        session, doc = mock_db_session
        file_path = fixtures_dir / "test_merged_data.xlsx"

        if not file_path.exists():
            pytest.skip("test_merged_data.xlsx not available")

        doc.file_path = str(file_path)
        doc.filename = "test_merged_data.xlsx"

        pipeline = ExtractionPipeline(db_session=session)
        summary = pipeline.process_document(doc.id)

        assert isinstance(summary, ExtractionSummary)
        assert summary.records_extracted > 0
        assert doc.status == "completed"

    def test_pipeline_with_multilevel_headers(self, fixtures_dir, mock_db_session):
        """Test pipeline handles multi-level headers correctly through all passes."""
        session, doc = mock_db_session
        file_path = fixtures_dir / "test_two_level_headers.xlsx"

        if not file_path.exists():
            pytest.skip("test_two_level_headers.xlsx not available")

        doc.file_path = str(file_path)
        doc.filename = "test_two_level_headers.xlsx"

        pipeline = ExtractionPipeline(db_session=session)
        summary = pipeline.process_document(doc.id)

        assert isinstance(summary, ExtractionSummary)
        assert summary.tables_detected >= 1
        assert doc.status == "completed"


class TestPipelineResilienceIntegration:
    """Test pipeline resilience to partial failures."""

    @pytest.fixture
    def fixtures_dir(self):
        """Get path to test fixtures directory."""
        from pathlib import Path

        return Path(__file__).parent.parent / "fixtures" / "sample_excel"

    @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()

        # Mock document
        from src.db import models as db_models

        doc = Mock(spec=db_models.Document)
        doc.id = uuid4()
        doc.filename = "test_standard_table.xlsx"
        doc.file_path = None
        doc.status = "pending"
        doc.sheet_count = None
        doc.record_count = None
        doc.processed_at = None
        doc.error_message = None
        doc.api_key_id = uuid4()

        session.scalar = Mock(return_value=doc)

        # Mock query chain - empty results
        mock_query = Mock()
        session.query = Mock(return_value=mock_query)
        mock_query.filter = Mock(return_value=mock_query)
        mock_query.order_by = Mock(return_value=mock_query)
        mock_query.all = Mock(return_value=[])
        mock_query.first = Mock(return_value=None)

        return session, doc

    def test_pipeline_completes_when_no_symbols_found(self, fixtures_dir, mock_db_session):
        """Test pipeline completes successfully when no symbol dictionaries are found."""
        session, doc = mock_db_session
        file_path = fixtures_dir / "test_standard_table.xlsx"
        doc.file_path = str(file_path)

        pipeline = ExtractionPipeline(db_session=session)
        summary = pipeline.process_document(doc.id)

        # Pipeline should complete even without symbols
        assert isinstance(summary, ExtractionSummary)
        assert doc.status == "completed"
        assert summary.symbols_detected == 0
        assert summary.symbols_resolved == 0

    def test_pipeline_completes_when_no_cross_refs_found(self, fixtures_dir, mock_db_session):
        """Test pipeline completes successfully when no cross-references are found."""
        session, doc = mock_db_session
        file_path = fixtures_dir / "test_standard_table.xlsx"
        doc.file_path = str(file_path)

        pipeline = ExtractionPipeline(db_session=session)
        summary = pipeline.process_document(doc.id)

        # Pipeline should complete even without cross-refs
        assert isinstance(summary, ExtractionSummary)
        assert doc.status == "completed"
        assert summary.cross_references_found == 0
        assert summary.cross_references_resolved == 0


class TestExtractionSummaryIntegration:
    """Test ExtractionSummary properly aggregates all pass results."""

    def test_summary_aggregates_warnings_from_all_passes(self):
        """Test that warnings from all passes are aggregated in summary."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=5000,
            warnings=[
                "Symbol detection: Unknown sheet format",
                "Table extraction: Skipped empty table",
                "Symbol resolution: Unknown symbol '△'",
                "Cross-ref: Unresolved reference 'See Sheet X'",
            ],
            symbols_detected=10,
            symbols_resolved=50,
            unresolved_symbols=2,
            cross_references_found=15,
            cross_references_resolved=12,
            cross_references_unresolved=3,
        )

        # All warnings should be preserved
        assert len(summary.warnings) == 4
        assert any("Symbol detection" in w for w in summary.warnings)
        assert any("Table extraction" in w for w in summary.warnings)
        assert any("Symbol resolution" in w for w in summary.warnings)
        assert any("Cross-ref" in w for w in summary.warnings)

    def test_summary_to_dict_complete(self):
        """Test that to_dict includes all statistics from all passes."""
        doc_id = uuid4()
        summary = ExtractionSummary(
            document_id=doc_id,
            sheets_processed=3,
            tables_detected=5,
            records_extracted=100,
            duration_ms=5000,
            warnings=[],
            symbols_detected=10,
            symbols_resolved=50,
            unresolved_symbols=2,
            cross_references_found=15,
            cross_references_resolved=12,
            cross_references_unresolved=3,
        )

        d = summary.to_dict()

        # Verify all fields are present
        expected_fields = [
            "document_id",
            "sheets_processed",
            "tables_detected",
            "records_extracted",
            "duration_ms",
            "warnings",
            "symbols_detected",
            "symbols_resolved",
            "unresolved_symbols",
            "cross_references_found",
            "cross_references_resolved",
            "cross_references_unresolved",
        ]

        for field in expected_fields:
            assert field in d, f"Missing field: {field}"

        # Verify values are correct
        assert d["symbols_detected"] == 10
        assert d["symbols_resolved"] == 50
        assert d["unresolved_symbols"] == 2
        assert d["cross_references_found"] == 15
        assert d["cross_references_resolved"] == 12
        assert d["cross_references_unresolved"] == 3
