"""Unit tests for HeaderExtractor."""

import pytest
from openpyxl import Workbook
from openpyxl.cell import MergedCell

from src.extraction.header_extractor import HeaderExtractor
from src.extraction.models import TableBoundary, HeaderInfo
from src.core.exceptions import ExtractionError


@pytest.fixture
def extractor():
    """Create a HeaderExtractor instance."""
    return HeaderExtractor()


class TestHeaderInfoModel:
    """Test HeaderInfo dataclass validation."""

    def test_valid_header_info(self):
        """Should create valid HeaderInfo with correct fields."""
        header = HeaderInfo(
            display="Sales > Q1",
            levels=["Sales", "Q1"],
            depth=2,
            column_letter="B",
        )

        assert header.display == "Sales > Q1"
        assert header.levels == ["Sales", "Q1"]
        assert header.depth == 2
        assert header.column_letter == "B"

    def test_invalid_depth_too_low(self):
        """Should raise ValueError for depth < 1."""
        with pytest.raises(ValueError, match="Must be between 1 and 4"):
            HeaderInfo(
                display="Product",
                levels=["Product"],
                depth=0,
                column_letter="A",
            )

    def test_invalid_depth_too_high(self):
        """Should raise ValueError for depth > 4."""
        with pytest.raises(ValueError, match="Must be between 1 and 4"):
            HeaderInfo(
                display="A > B > C > D > E",
                levels=["A", "B", "C", "D", "E"],
                depth=5,
                column_letter="A",
            )

    def test_empty_levels(self):
        """Should raise ValueError for empty levels."""
        with pytest.raises(ValueError, match="levels cannot be empty"):
            HeaderInfo(
                display="",
                levels=[],
                depth=1,
                column_letter="A",
            )

    def test_depth_mismatch(self):
        """Should raise ValueError when depth doesn't match levels count."""
        with pytest.raises(ValueError, match="must match number of levels"):
            HeaderInfo(
                display="Sales > Q1",
                levels=["Sales", "Q1"],
                depth=3,
                column_letter="B",
            )

    def test_empty_column_letter(self):
        """Should raise ValueError for empty column letter."""
        with pytest.raises(ValueError, match="Column letter cannot be empty"):
            HeaderInfo(
                display="Product",
                levels=["Product"],
                depth=1,
                column_letter="",
            )

    def test_to_dict(self):
        """Should convert HeaderInfo to dictionary."""
        header = HeaderInfo(
            display="Inventory > Current",
            levels=["Inventory", "Current"],
            depth=2,
            column_letter="D",
        )

        result = header.to_dict()

        assert result == {
            "display": "Inventory > Current",
            "levels": ["Inventory", "Current"],
            "depth": 2,
            "column_letter": "D",
        }


class TestSingleRowHeader:
    """Test single-row header extraction."""

    def test_simple_single_row_header(self, extractor):
        """Should extract single-row header correctly."""
        wb = Workbook()
        sheet = wb.active

        # Create single-row header
        sheet["A1"] = "Product"
        sheet["B1"] = "Price"
        sheet["C1"] = "Quantity"
        sheet["D1"] = "Total"

        # Add data row
        sheet["A2"] = "Widget"
        sheet["B2"] = 10.50
        sheet["C2"] = 5
        sheet["D2"] = 52.50

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=2, start_col="A", end_col="D"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 4
        assert headers["A"].display == "Product"
        assert headers["A"].depth == 1
        assert headers["B"].display == "Price"
        assert headers["C"].display == "Quantity"
        assert headers["D"].display == "Total"


class TestTwoLevelHeader:
    """Test two-level header extraction."""

    def test_simple_two_level_header(self, extractor):
        """Should extract two-level headers with parent > child format."""
        wb = Workbook()
        sheet = wb.active

        # Create two-level header
        sheet["A1"] = "Product"
        sheet["A2"] = "Name"
        sheet["B1"] = "Sales"
        sheet["B2"] = "Q1"
        sheet["C1"] = "Sales"
        sheet["C2"] = "Q2"
        sheet["D1"] = "Total"
        sheet["D2"] = "Amount"

        # Add data row
        sheet["A3"] = "Widget"
        sheet["B3"] = 100
        sheet["C3"] = 150
        sheet["D3"] = 250

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="D"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 4
        assert headers["A"].display == "Product > Name"
        assert headers["A"].depth == 2
        assert headers["A"].levels == ["Product", "Name"]
        assert headers["B"].display == "Sales > Q1"
        assert headers["B"].levels == ["Sales", "Q1"]
        assert headers["C"].display == "Sales > Q2"
        assert headers["D"].display == "Total > Amount"

    def test_two_level_with_merged_parent(self, extractor):
        """Should handle merged parent cell spanning multiple children."""
        wb = Workbook()
        sheet = wb.active

        # Create header with merged parent
        sheet["A1"] = "Product"
        sheet["A2"] = "Name"

        # Merge B1:C1 for "Sales" parent
        sheet["B1"] = "Sales"
        sheet.merge_cells("B1:C1")
        sheet["B2"] = "Q1"
        sheet["C2"] = "Q2"

        # Add data row
        sheet["A3"] = "Widget"
        sheet["B3"] = 100
        sheet["C3"] = 150

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="C"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 3
        assert headers["B"].display == "Sales > Q1"
        assert headers["B"].levels == ["Sales", "Q1"]
        assert headers["C"].display == "Sales > Q2"
        assert headers["C"].levels == ["Sales", "Q2"]


class TestThreeLevelHeader:
    """Test three-level header extraction."""

    def test_three_level_header(self, extractor):
        """Should extract three-level headers correctly."""
        wb = Workbook()
        sheet = wb.active

        # Create three-level header
        sheet["A1"] = "Product"
        sheet["A2"] = "Info"
        sheet["A3"] = "Name"

        sheet["B1"] = "Sales"
        sheet["B2"] = "Q1"
        sheet["B3"] = "Amount"

        sheet["C1"] = "Sales"
        sheet["C2"] = "Q1"
        sheet["C3"] = "Units"

        # Add data row
        sheet["A4"] = "Widget"
        sheet["B4"] = 1000
        sheet["C4"] = 50

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=4, start_col="A", end_col="C"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 3
        assert headers["A"].display == "Product > Info > Name"
        assert headers["A"].depth == 3
        assert headers["A"].levels == ["Product", "Info", "Name"]
        assert headers["B"].display == "Sales > Q1 > Amount"
        assert headers["C"].display == "Sales > Q1 > Units"


class TestFourLevelHeader:
    """Test four-level header extraction (maximum depth)."""

    def test_four_level_header(self, extractor):
        """Should extract four-level headers (max depth)."""
        wb = Workbook()
        sheet = wb.active

        # Create four-level header
        sheet["A1"] = "Company"
        sheet["A2"] = "Division"
        sheet["A3"] = "Department"
        sheet["A4"] = "Team"

        sheet["B1"] = "Metrics"
        sheet["B2"] = "Sales"
        sheet["B3"] = "Q1"
        sheet["B4"] = "Revenue"

        # Add data row
        sheet["A5"] = "Engineering"
        sheet["B5"] = 100000

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=5, start_col="A", end_col="B"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 2
        assert headers["A"].display == "Company > Division > Department > Team"
        assert headers["A"].depth == 4
        assert headers["B"].display == "Metrics > Sales > Q1 > Revenue"


class TestMixedDepthColumns:
    """Test headers with mixed depth columns."""

    def test_mixed_depth_columns(self, extractor):
        """Should handle columns with different depths."""
        wb = Workbook()
        sheet = wb.active

        # Create mixed depth header
        # Column A: 1 level
        sheet["A1"] = "Product"
        sheet["A2"] = "Product"  # Repeat to match depth

        # Columns B, C: 2 levels
        sheet["B1"] = "Sales"
        sheet["B2"] = "Q1"
        sheet["C1"] = "Sales"
        sheet["C2"] = "Q2"

        # Add data row
        sheet["A3"] = "Widget"
        sheet["B3"] = 100
        sheet["C3"] = 150

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="C"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 3
        # Column A has repeated value, so depth = 2
        assert headers["A"].display == "Product > Product"
        assert headers["B"].display == "Sales > Q1"
        assert headers["C"].display == "Sales > Q2"


class TestEmptyCellPropagation:
    """Test empty cell propagation in headers."""

    def test_empty_cell_propagates_parent(self, extractor):
        """Should propagate parent value when child cell is empty."""
        wb = Workbook()
        sheet = wb.active

        # Create header with empty child cells
        sheet["A1"] = "Sales"
        sheet["A2"] = ""  # Empty - should propagate "Sales"

        sheet["B1"] = "Sales"
        sheet["B2"] = "Q1"

        # Add data row
        sheet["A3"] = 100
        sheet["B3"] = 100

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="B"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 2
        # Empty B2 propagates "Sales"
        assert headers["A"].display == "Sales > Sales"
        assert headers["B"].display == "Sales > Q1"


class TestEdgeCases:
    """Test edge cases and error conditions."""

    def test_all_empty_header_row(self, extractor):
        """Should handle case where header row is empty."""
        wb = Workbook()
        sheet = wb.active

        # Empty header row
        sheet["A1"] = ""
        sheet["B1"] = ""
        sheet["C1"] = ""

        # Data row
        sheet["A2"] = "data1"
        sheet["B2"] = "data2"
        sheet["C2"] = "data3"

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=2, start_col="A", end_col="C"
        )

        headers = extractor.extract_headers(sheet, boundary)

        # All columns have empty headers, should return empty dict
        assert len(headers) == 0

    def test_trailing_empty_columns(self, extractor):
        """Should exclude trailing empty columns from results."""
        wb = Workbook()
        sheet = wb.active

        # Header with trailing empty columns
        sheet["A1"] = "Product"
        sheet["B1"] = "Price"
        sheet["C1"] = ""  # Empty
        sheet["D1"] = ""  # Empty

        # Data
        sheet["A2"] = "Widget"
        sheet["B2"] = 10.50

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=2, start_col="A", end_col="D"
        )

        headers = extractor.extract_headers(sheet, boundary)

        # Only non-empty header columns
        assert len(headers) == 2
        assert "A" in headers
        assert "B" in headers
        assert "C" not in headers
        assert "D" not in headers


class TestHeaderDepthDetection:
    """Test header depth detection algorithm."""

    def test_detect_single_row_header(self, extractor):
        """Should detect depth=1 for single row header."""
        wb = Workbook()
        sheet = wb.active

        # Single row header (text)
        sheet["A1"] = "Product"
        sheet["B1"] = "Price"

        # Data row (numbers)
        sheet["A2"] = "Widget"
        sheet["B2"] = 10.50

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=2, start_col="A", end_col="B"
        )

        depth = extractor._detect_header_depth(sheet, boundary)
        assert depth == 1

    def test_detect_two_row_header(self, extractor):
        """Should detect depth=2 for two row header."""
        wb = Workbook()
        sheet = wb.active

        # Two row header (text)
        sheet["A1"] = "Sales"
        sheet["B1"] = "Sales"
        sheet["A2"] = "Q1"
        sheet["B2"] = "Q2"

        # Data row (numbers)
        sheet["A3"] = 100
        sheet["B3"] = 150

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="B"
        )

        depth = extractor._detect_header_depth(sheet, boundary)
        assert depth == 2

    def test_detect_stops_at_data_row(self, extractor):
        """Should stop counting when hitting data row."""
        wb = Workbook()
        sheet = wb.active

        # One header row
        sheet["A1"] = "Product"
        sheet["B1"] = "Price"

        # Data rows (should stop here)
        sheet["A2"] = "Widget"
        sheet["B2"] = 10.50
        sheet["A3"] = "Gadget"
        sheet["B3"] = 20.00

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=3, start_col="A", end_col="B"
        )

        depth = extractor._detect_header_depth(sheet, boundary)
        assert depth == 1  # Only first row is header


class TestIntegrationWithTableBoundary:
    """Test integration with TableBoundary from table detection."""

    def test_uses_boundary_start_row(self, extractor):
        """Should use boundary.start_row to locate headers."""
        wb = Workbook()
        sheet = wb.active

        # Empty rows
        sheet["A1"] = ""
        sheet["A2"] = ""

        # Header starts at row 3
        sheet["A3"] = "Product"
        sheet["B3"] = "Price"

        # Data
        sheet["A4"] = "Widget"
        sheet["B4"] = 10.50

        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=3, end_row=4, start_col="A", end_col="B"
        )

        headers = extractor.extract_headers(sheet, boundary)

        assert len(headers) == 2
        assert headers["A"].display == "Product"
        assert headers["B"].display == "Price"

    def test_respects_column_range(self, extractor):
        """Should only extract headers within boundary column range."""
        wb = Workbook()
        sheet = wb.active

        # Headers in columns A-E
        sheet["A1"] = "Col A"
        sheet["B1"] = "Col B"
        sheet["C1"] = "Col C"
        sheet["D1"] = "Col D"
        sheet["E1"] = "Col E"

        # Boundary only covers B-D
        boundary = TableBoundary(
            sheet="Sheet", table_id=1, start_row=1, end_row=2, start_col="B", end_col="D"
        )

        headers = extractor.extract_headers(sheet, boundary)

        # Should only get B, C, D
        assert len(headers) == 3
        assert "A" not in headers
        assert "B" in headers
        assert "C" in headers
        assert "D" in headers
        assert "E" not in headers
