"""Merged cell resolution for Excel data rows."""

import structlog
from openpyxl.worksheet.worksheet import Worksheet
from openpyxl.utils import column_index_from_string, get_column_letter

from src.extraction.models import TableBoundary, MergeMetadata
from src.core.exceptions import ExtractionError

logger = structlog.get_logger(__name__)


class MergedCellResolver:
    """Resolves merged cells in Excel data rows by propagating values.

    Detects merged cell ranges within a table boundary and propagates
    the top-left (master) cell value to all cells in the merged range.
    """

    def resolve_merged_cells(
        self, sheet: Worksheet, table_boundary: TableBoundary
    ) -> tuple[list[list[str]], list[MergeMetadata]]:
        """Resolve merged cells within a table boundary.

        Args:
            sheet: openpyxl Worksheet object
            table_boundary: TableBoundary defining the region to process

        Returns:
            Tuple of (resolved_grid, metadata):
                - resolved_grid: 2D list where resolved[row_idx][col_idx] = cell value
                - metadata: List of MergeMetadata for propagated cells

        Raises:
            ExtractionError: If merge resolution fails
        """
        try:
            start_time = logger.bind(
                sheet_name=sheet.title, table_id=table_boundary.table_id
            )

            # Get all merged ranges in sheet
            all_merged_ranges = list(sheet.merged_cells.ranges)

            # Filter to ranges within boundary
            relevant_ranges = self._filter_merged_ranges(
                all_merged_ranges, table_boundary
            )

            logger.info(
                "merged_ranges_detected",
                sheet_name=sheet.title,
                table_id=table_boundary.table_id,
                total_ranges=len(all_merged_ranges),
                relevant_ranges=len(relevant_ranges),
            )

            # Build initial grid from sheet data
            resolved_grid = self._build_initial_grid(sheet, table_boundary)

            # Propagate merged cell values and track metadata
            metadata = self._propagate_merged_values(
                sheet, table_boundary, relevant_ranges, resolved_grid
            )

            logger.info(
                "merged_cells_resolved",
                sheet_name=sheet.title,
                table_id=table_boundary.table_id,
                propagated_cells=len(metadata),
            )

            return resolved_grid, metadata

        except Exception as e:
            logger.error(
                "merge_resolution_failed",
                sheet_name=sheet.title,
                table_id=table_boundary.table_id,
                error=str(e),
            )
            raise ExtractionError(
                code="MERGE_RESOLUTION_FAILED",
                message=f"Failed to resolve merged cells in table {table_boundary.table_id}",
                details={"sheet": sheet.title, "error": str(e)},
            )

    def _filter_merged_ranges(
        self, merged_ranges: list, table_boundary: TableBoundary
    ) -> list:
        """Filter merged ranges to only those intersecting the table boundary.

        Args:
            merged_ranges: All merged cell ranges in the sheet
            table_boundary: TableBoundary to check intersection with

        Returns:
            List of merged ranges that intersect with the boundary
        """
        start_col_idx = column_index_from_string(table_boundary.start_col)
        end_col_idx = column_index_from_string(table_boundary.end_col)

        relevant_ranges = []

        for merged_range in merged_ranges:
            # Check if range intersects with table boundary
            if self._intersects_boundary(
                merged_range,
                table_boundary.start_row,
                table_boundary.end_row,
                start_col_idx,
                end_col_idx,
            ):
                relevant_ranges.append(merged_range)

        return relevant_ranges

    def _intersects_boundary(
        self,
        merged_range,
        start_row: int,
        end_row: int,
        start_col_idx: int,
        end_col_idx: int,
    ) -> bool:
        """Check if a merged range intersects with the table boundary.

        Args:
            merged_range: openpyxl MergedCellRange object
            start_row: Table start row (1-indexed)
            end_row: Table end row (1-indexed)
            start_col_idx: Table start column index (1-indexed)
            end_col_idx: Table end column index (1-indexed)

        Returns:
            True if the merged range intersects with the boundary
        """
        # Get range boundaries
        range_start_row = merged_range.min_row
        range_end_row = merged_range.max_row
        range_start_col = merged_range.min_col
        range_end_col = merged_range.max_col

        # Check for intersection
        row_intersects = not (range_end_row < start_row or range_start_row > end_row)
        col_intersects = not (
            range_end_col < start_col_idx or range_start_col > end_col_idx
        )

        return row_intersects and col_intersects

    def _build_initial_grid(
        self, sheet: Worksheet, table_boundary: TableBoundary
    ) -> list[list[str]]:
        """Build initial 2D grid with all cell values.

        Args:
            sheet: openpyxl Worksheet object
            table_boundary: TableBoundary defining the region

        Returns:
            2D list with initial cell values (before merge propagation)
        """
        start_col_idx = column_index_from_string(table_boundary.start_col)
        end_col_idx = column_index_from_string(table_boundary.end_col)

        num_rows = table_boundary.end_row - table_boundary.start_row + 1
        num_cols = end_col_idx - start_col_idx + 1

        resolved_grid = []

        for row_offset in range(num_rows):
            row_idx = table_boundary.start_row + row_offset
            row_values = []

            for col_offset in range(num_cols):
                col_idx = start_col_idx + col_offset
                cell = sheet.cell(row=row_idx, column=col_idx)

                # Get cell value, handling None
                value = cell.value if cell.value is not None else ""
                row_values.append(str(value).strip() if value else "")

            resolved_grid.append(row_values)

        return resolved_grid

    def _propagate_merged_values(
        self,
        sheet: Worksheet,
        table_boundary: TableBoundary,
        merged_ranges: list,
        resolved_grid: list[list[str]],
    ) -> list[MergeMetadata]:
        """Propagate merged cell values to all cells in each range.

        Args:
            sheet: openpyxl Worksheet object
            table_boundary: TableBoundary defining the region
            merged_ranges: List of merged ranges to process
            resolved_grid: 2D grid to update with propagated values

        Returns:
            List of MergeMetadata for all propagated cells
        """
        start_col_idx = column_index_from_string(table_boundary.start_col)
        metadata_list = []

        for merged_range in merged_ranges:
            # Get master cell (top-left) value
            master_cell = sheet.cell(
                row=merged_range.min_row, column=merged_range.min_col
            )
            master_value = (
                str(master_cell.value).strip()
                if master_cell.value is not None
                else ""
            )

            # Get range coordinate string
            range_coord = merged_range.coord

            # Check for overlapping ranges
            if self._has_overlapping_values(
                sheet, merged_range, master_value, resolved_grid, table_boundary
            ):
                logger.warning(
                    "overlapping_merged_range",
                    sheet_name=sheet.title,
                    table_id=table_boundary.table_id,
                    range=range_coord,
                    message="Using first encountered value",
                )

            # Propagate value to all cells in range
            for row in range(merged_range.min_row, merged_range.max_row + 1):
                for col in range(merged_range.min_col, merged_range.max_col + 1):
                    # Skip if outside table boundary
                    if (
                        row < table_boundary.start_row
                        or row > table_boundary.end_row
                        or col < column_index_from_string(table_boundary.start_col)
                        or col > column_index_from_string(table_boundary.end_col)
                    ):
                        continue

                    # Convert to grid coordinates
                    grid_row = row - table_boundary.start_row
                    grid_col = col - start_col_idx

                    # Update grid with master value
                    resolved_grid[grid_row][grid_col] = master_value

                    # Track metadata (skip master cell itself)
                    if row != merged_range.min_row or col != merged_range.min_col:
                        cell_coord = f"{get_column_letter(col)}{row}"
                        metadata = MergeMetadata(
                            cell_coordinate=cell_coord,
                            merged_from_range=range_coord,
                            original_value=master_value,
                        )
                        metadata_list.append(metadata)

        return metadata_list

    def _has_overlapping_values(
        self,
        sheet: Worksheet,
        merged_range,
        expected_value: str,
        resolved_grid: list[list[str]],
        table_boundary: TableBoundary,
    ) -> bool:
        """Check if merged range has overlapping or conflicting values.

        Args:
            sheet: openpyxl Worksheet object
            merged_range: Merged range to check
            expected_value: Expected master cell value
            resolved_grid: Current grid state
            table_boundary: TableBoundary defining the region

        Returns:
            True if overlapping values detected
        """
        start_col_idx = column_index_from_string(table_boundary.start_col)

        for row in range(merged_range.min_row, merged_range.max_row + 1):
            for col in range(merged_range.min_col, merged_range.max_col + 1):
                # Skip if outside boundary
                if (
                    row < table_boundary.start_row
                    or row > table_boundary.end_row
                    or col < column_index_from_string(table_boundary.start_col)
                    or col > column_index_from_string(table_boundary.end_col)
                ):
                    continue

                grid_row = row - table_boundary.start_row
                grid_col = col - start_col_idx

                current_value = resolved_grid[grid_row][grid_col]

                # Check if cell already has a different non-empty value
                if current_value and current_value != expected_value:
                    return True

        return False
