from __future__ import annotations

import base64
import copy
import json
import uuid
from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path
from typing import Any


class ElementType(str, Enum):
    TEXT = "text"
    IMAGE = "image"
    TABLE = "table"


@dataclass
class FontSettings:
    family: str = "Arial"
    size_pt: float = 12.0
    bold: bool = False
    italic: bool = False
    underline: bool = False

    def to_dict(self) -> dict[str, Any]:
        return {
            "family": self.family,
            "size_pt": self.size_pt,
            "bold": self.bold,
            "italic": self.italic,
            "underline": self.underline,
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any]) -> FontSettings:
        return cls(
            family=str(data.get("family", "Arial")),
            size_pt=float(data.get("size_pt", 12.0)),
            bold=bool(data.get("bold", False)),
            italic=bool(data.get("italic", False)),
            underline=bool(data.get("underline", False)),
        )


@dataclass
class TableData:
    rows: list[list[str]] = field(default_factory=list)
    column_widths: list[float] = field(default_factory=list)
    row_heights: list[float] = field(default_factory=list)

    def to_dict(self) -> dict[str, Any]:
        return {
            "rows": self.rows,
            "column_widths": self.column_widths,
            "row_heights": self.row_heights,
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any]) -> TableData:
        return cls(
            rows=[
                [str(value) for value in row]
                for row in data.get("rows", [])
            ],
            column_widths=[
                float(value) for value in data.get("column_widths", [])
            ],
            row_heights=[
                float(value) for value in data.get("row_heights", [])
            ],
        )


@dataclass
class DocumentElement:
    element_type: ElementType
    x: float
    y: float
    width: float
    height: float
    page_number: int
    id: str = field(default_factory=lambda: uuid.uuid4().hex)
    text: str = ""
    image_png_b64: str = ""
    table: TableData | None = None
    processing: bool = False
    error: str = ""
    table_fallback_image: bool = False

    def __post_init__(self) -> None:
        self.clamp()

    def clamp(self) -> None:
        self.x = min(max(float(self.x), 0.0), 1.0)
        self.y = min(max(float(self.y), 0.0), 1.0)
        self.width = min(max(float(self.width), 0.001), 1.0 - self.x)
        self.height = min(max(float(self.height), 0.001), 1.0 - self.y)

    def rect(self) -> tuple[float, float, float, float]:
        return self.x, self.y, self.width, self.height

    def to_dict(self) -> dict[str, Any]:
        return {
            "id": self.id,
            "element_type": self.element_type.value,
            "x": self.x,
            "y": self.y,
            "width": self.width,
            "height": self.height,
            "page_number": self.page_number,
            "text": self.text,
            "image_png_b64": self.image_png_b64,
            "table": self.table.to_dict() if self.table else None,
            "processing": self.processing,
            "error": self.error,
            "table_fallback_image": self.table_fallback_image,
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any]) -> DocumentElement:
        table = data.get("table")
        return cls(
            id=str(data.get("id") or uuid.uuid4().hex),
            element_type=ElementType(str(data["element_type"])),
            x=float(data["x"]),
            y=float(data["y"]),
            width=float(data["width"]),
            height=float(data["height"]),
            page_number=int(data["page_number"]),
            text=str(data.get("text", "")),
            image_png_b64=str(data.get("image_png_b64", "")),
            table=TableData.from_dict(table) if table else None,
            processing=bool(data.get("processing", False)),
            error=str(data.get("error", "")),
            table_fallback_image=bool(
                data.get("table_fallback_image", False)
            ),
        )


@dataclass
class PageState:
    page_number: int
    width_pt: float
    height_pt: float
    processed: bool = False
    elements: list[DocumentElement] = field(default_factory=list)

    @property
    def orientation(self) -> str:
        return "landscape" if self.width_pt > self.height_pt else "portrait"

    def to_dict(self) -> dict[str, Any]:
        return {
            "page_number": self.page_number,
            "width_pt": self.width_pt,
            "height_pt": self.height_pt,
            "processed": self.processed,
            "elements": [element.to_dict() for element in self.elements],
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any]) -> PageState:
        return cls(
            page_number=int(data["page_number"]),
            width_pt=float(data["width_pt"]),
            height_pt=float(data["height_pt"]),
            processed=bool(data.get("processed", False)),
            elements=[
                DocumentElement.from_dict(element)
                for element in data.get("elements", [])
            ],
        )


@dataclass
class DocumentModel:
    source_pdf: str = ""
    output_docx: str = ""
    current_page: int = 1
    pages: list[PageState] = field(default_factory=list)
    font: FontSettings = field(default_factory=FontSettings)
    ocr_language: str = "eng"
    render_dpi: int = 180
    preprocess: bool = True
    tesseract_path: str = ""

    def page(self, page_number: int) -> PageState:
        if page_number < 1 or page_number > len(self.pages):
            raise IndexError(f"Invalid page number: {page_number}")
        return self.pages[page_number - 1]

    def element(self, element_id: str) -> DocumentElement | None:
        for page in self.pages:
            for element in page.elements:
                if element.id == element_id:
                    return element
        return None

    def clone(self) -> DocumentModel:
        return copy.deepcopy(self)

    def to_dict(self) -> dict[str, Any]:
        return {
            "version": 1,
            "source_pdf": self.source_pdf,
            "output_docx": self.output_docx,
            "current_page": self.current_page,
            "font": self.font.to_dict(),
            "ocr_language": self.ocr_language,
            "render_dpi": self.render_dpi,
            "preprocess": self.preprocess,
            "tesseract_path": self.tesseract_path,
            "pages": [page.to_dict() for page in self.pages],
        }

    @classmethod
    def from_dict(cls, data: dict[str, Any]) -> DocumentModel:
        return cls(
            source_pdf=str(data.get("source_pdf", "")),
            output_docx=str(data.get("output_docx", "")),
            current_page=int(data.get("current_page", 1)),
            font=FontSettings.from_dict(data.get("font", {})),
            ocr_language=str(data.get("ocr_language", "eng")),
            render_dpi=int(data.get("render_dpi", 180)),
            preprocess=bool(data.get("preprocess", True)),
            tesseract_path=str(data.get("tesseract_path", "")),
            pages=[
                PageState.from_dict(page)
                for page in data.get("pages", [])
            ],
        )

    def save_project(self, path: str | Path) -> None:
        destination = Path(path)
        destination.parent.mkdir(parents=True, exist_ok=True)

        payload = json.dumps(
            self.to_dict(),
            indent=2,
            ensure_ascii=False,
        )
        destination.write_text(payload, encoding="utf-8")

    @classmethod
    def load_project(cls, path: str | Path) -> DocumentModel:
        payload = Path(path).read_text(encoding="utf-8")
        return cls.from_dict(json.loads(payload))

    @staticmethod
    def image_to_b64(data: bytes) -> str:
        return base64.b64encode(data).decode("ascii")

    @staticmethod
    def b64_to_image(value: str) -> bytes:
        return base64.b64decode(value.encode("ascii"))