"""
Header consolidation post-processor.

Operates on a DoclingDocument (the canonical Parser output) and returns
TableMetadata for each table. Pulled out of the parser to keep that layer
engine-agnostic — the consolidation logic is independent of how the
DoclingDocument was produced.
"""

import json
from datetime import datetime
from pathlib import Path
from typing import Any, List, Optional, Tuple

from docling_core.types.doc.document import DoclingDocument

from quber.core.models import TableMetadata


class HeaderConsolidator:
    def __init__(
        self,
        separator: str = " ^ ",
        max_lookback: int = 10,
        enable_consolidation: bool = True,
        debug: bool = False,
    ) -> None:
        self.separator = separator
        self.max_lookback = max_lookback
        self.enable_consolidation = enable_consolidation
        self.debug = debug

    def consolidate(self, document: DoclingDocument) -> List[TableMetadata]:
        tables_metadata: List[TableMetadata] = []
        table_index = 0

        elements_with_depth = list(document.iterate_items())
        elements_list = [item for item, _ in elements_with_depth]
        element_to_index = {id(elem): i for i, elem in enumerate(elements_list)}

        for table in document.tables:
            page_num = (
                table.prov[0].page_no
                if (hasattr(table, "prov") and table.prov and hasattr(table.prov[0], "page_no"))
                else 1
            )

            table_idx = element_to_index.get(id(table), -1)

            headers: List[str] = []
            descriptive_texts: List[str] = []
            if table_idx > 0:
                headers, descriptive_texts = self.find_preceding_headers(elements_list, table_idx)

            consolidated_header: Optional[str] = None
            if self.enable_consolidation and headers:
                consolidated_header = self.consolidate_headers(headers)

            rows = 0
            cols = 0
            first_cell: Optional[str] = None
            if hasattr(table, "data") and table.data:
                rows = table.data.num_rows
                cols = table.data.num_cols
                if rows > 0 and cols > 0:
                    if hasattr(table.data, "grid") and table.data.grid:
                        first_row = table.data.grid[0] if len(table.data.grid) > 0 else []
                        if first_row and len(first_row) > 0:
                            first_cell = str(first_row[0]) if first_row[0] else None

            preceding_text = self.get_preceding_text(elements_list, table_idx, limit=3)

            table_md = table.export_to_markdown(document) if hasattr(table, "export_to_markdown") else None

            metadata = TableMetadata(
                table_index=table_index,
                page_number=page_num,
                headers=headers,
                consolidated_header=consolidated_header,
                descriptive_text=descriptive_texts[0] if descriptive_texts else None,
                table_markdown=table_md,
                preceding_text=preceding_text,
                rows=rows,
                cols=cols,
                first_cell_content=first_cell,
            )
            tables_metadata.append(metadata)
            table_index += 1

            if self.debug:
                self.print_debug_info(metadata)

        return tables_metadata

    def find_preceding_headers(self, elements: List[Any], table_idx: int) -> Tuple[List[str], List[str]]:
        headers: List[str] = []
        descriptive_texts: List[str] = []
        current_level: Optional[int] = None
        consecutive = True

        for i in range(min(self.max_lookback, table_idx)):
            idx = table_idx - i - 1
            elem = elements[idx]
            elem_type = type(elem).__name__

            if elem_type in ["TableItem", "PictureItem"]:
                break

            if elem_type == "ListItem" and headers:
                break

            if not hasattr(elem, "text") or not elem.text:
                continue

            text = elem.text.strip()
            if not text:
                continue

            is_header = hasattr(elem, "level") and elem.level is not None
            if is_header:
                level = elem.level
                if current_level is None:
                    current_level = level
                    headers.insert(0, text)
                elif level == current_level and consecutive:
                    headers.insert(0, text)
                else:
                    break
            else:
                if not headers and len(descriptive_texts) == 0:
                    descriptive_texts.insert(0, text)
                elif headers:
                    consecutive = False

        return headers, descriptive_texts

    def consolidate_headers(self, headers: List[str]) -> Optional[str]:
        if not headers:
            return None
        if len(headers) == 1:
            return headers[0]
        return self.separator.join(headers)

    def get_preceding_text(self, elements: List[Any], table_idx: int, limit: int = 3) -> Optional[str]:
        texts: List[str] = []
        count = 0

        for i in range(min(self.max_lookback, table_idx)):
            idx = table_idx - i - 1
            elem = elements[idx]
            elem_type = type(elem).__name__

            if elem_type == "TableItem":
                texts.insert(0, "[Table]")
                continue

            if hasattr(elem, "text") and elem.text:
                text = elem.text.strip()
                if text:
                    texts.insert(0, text)
                    count += 1
                    if count >= limit:
                        break

        return "\n".join(texts) if texts else None

    def print_debug_info(self, metadata: TableMetadata) -> None:
        print(f"\n=== Table {metadata.table_index + 1} (Page {metadata.page_number}) ===")
        if metadata.consolidated_header:
            print(f"Consolidated Header: {metadata.consolidated_header}")
        if metadata.headers:
            print(f"Individual Headers: {metadata.headers}")
        if metadata.descriptive_text:
            print(f"Descriptive Text: {metadata.descriptive_text[:100]}...")
        print(f"Size: {metadata.rows} rows x {metadata.cols} cols")
        if metadata.first_cell_content:
            print(f"First Cell: {metadata.first_cell_content[:50]}...")


def save_metadata(metadata_list: List[TableMetadata], output_path: str | Path) -> None:
    data = {
        "extraction_date": datetime.now().isoformat(),
        "total_tables": len(metadata_list),
        "tables": [
            {
                "table_index": meta.table_index,
                "page_number": meta.page_number,
                "consolidated_header": meta.consolidated_header,
                "headers": meta.headers,
                "descriptive_text": meta.descriptive_text,
                "preceding_text": meta.preceding_text,
                "dimensions": {"rows": meta.rows, "cols": meta.cols},
                "first_cell": meta.first_cell_content,
            }
            for meta in metadata_list
        ],
    }
    with open(output_path, "w", encoding="utf-8") as f:
        json.dump(data, f, indent=2, ensure_ascii=False)
