from __future__ import annotations

import io, os
import logging
from pathlib import Path
from typing import Any

import cv2
import fitz  # type: ignore[import-untyped]
import numpy as np
import pytesseract  # type: ignore[import-untyped]
from docx import Document
from docx.document import Document as DocumentObject

from docx.enum.section import WD_SECTION
from docx.enum.text import WD_ALIGN_PARAGRAPH

from docx.oxml import OxmlElement
from docx.oxml.ns import nsmap, qn

from docx.shared import Pt
from PIL import Image, ImageOps

from .models import (
    DocumentModel,
    ElementType,
    FontSettings,
    TableData,
)
from .utils import CoordinateMapper

nsmap.update( { "v": "urn:schemas-microsoft-com:vml", "w10": "urn:schemas-microsoft-com:office:word", } )

class PDFService:
    """Safe, lazy PDF access service.

    A new PyMuPDF document is opened for each operation. This avoids
    sharing a mutable PDF document between worker threads.
    """

    @staticmethod
    def validate(path: str) -> tuple[bool, str]:
        try:
            with fitz.open(path) as document:
                if document.needs_pass:
                    return False, "The PDF is password protected."
                if document.page_count <= 0:
                    return False, "The PDF contains no pages."
            return True, ""
        except Exception as exc:
            logger.exception("PDF validation failed")
            return False, str(exc)

    @staticmethod
    def page_count(path: str) -> int:
        with fitz.open(path) as document:
            if document.needs_pass:
                raise RuntimeError("The PDF is password protected.")
            return document.page_count

    @staticmethod
    def page_size(path: str, page_number: int) -> tuple[float, float]:
        with fitz.open(path) as document:
            page = document[page_number - 1]
            rect = page.rect
            return float(rect.width), float(rect.height)

    @staticmethod
    def render_page(
        path: str,
        page_number: int,
        dpi: int,
    ) -> bytes:
        with fitz.open(path) as document:
            page = document[page_number - 1]
            pixmap = page.get_pixmap(
                dpi=dpi,
                colorspace=fitz.csRGB,
                alpha=False,
            )
            return pixmap.tobytes("png")

    @staticmethod
    def render_region(
        path: str,
        page_number: int,
        normalized_rect: tuple[float, float, float, float],
        dpi: int,
    ) -> bytes:
        x, y, width, height = normalized_rect

        with fitz.open(path) as document:
            page = document[page_number - 1]
            rect = page.rect

            pdf_x, pdf_y, pdf_w, pdf_h = (
                CoordinateMapper.normalized_to_pdf(
                    x,
                    y,
                    width,
                    height,
                    rect.width,
                    rect.height,
                )
            )

            clip = fitz.Rect(
                pdf_x,
                pdf_y,
                pdf_x + pdf_w,
                pdf_y + pdf_h,
            )
            clip &= rect

            pixmap = page.get_pixmap(
                dpi=dpi,
                clip=clip,
                colorspace=fitz.csRGB,
                alpha=False,
            )
            return pixmap.tobytes("png")


from dotenv import load_dotenv
from google import genai
from google.genai import types

logger = logging.getLogger(__name__)

load_dotenv()


class OCRService:
    """Gemini API-based vision service for text extraction."""

    def __init__(
        self,
        executable: str = "",
        language: str = "eng",
        preprocess: bool = False,
    ) -> None:
        api_key = os.getenv("GEMINI_API_KEY")
        if not api_key:
            raise RuntimeError("GEMINI_API_KEY was not found in the environment or .env file.")

        self.client = genai.Client(api_key=api_key)
        self.model_name = "gemma-4-31b-it"
        self.language = language

    def recognize(
        self,
        image_bytes: bytes,
        psm: int = 6,
    ) -> str:
        try:
            response = self.client.models.generate_content(
                model=self.model_name,
                contents=[
                    types.Part.from_bytes(
                        data=image_bytes,
                        mime_type="image/png",
                    ),
                    (
                        "Extract all text, numbers, and structured content "
                        "from this image precisely as they appear. "
                        "Return only the extracted text without introductory or concluding remarks."
                    ),
                ],
            )
            return response.text.replace("\r\n", "\n").strip()
        except Exception as exc:
            logger.exception("Gemini API OCR failed")
            raise RuntimeError(f"Gemini OCR error: {exc}") from exc

    def languages(self) -> list[str]:
        return [self.language]

class TableService:
    """Gemini API-driven table recognizer."""

    def __init__(
        self,
        ocr_service: OCRService,
    ) -> None:
        self.ocr = ocr_service

    def recognize(
        self,
        image_bytes: bytes,
    ) -> TableData | None:
        image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
        width, height = image.size

        if width == 0 or height == 0:
            return None

        # Prompt Gemini specifically for CSV table output
        prompt = (
            "Analyze this table image. Extract all rows and columns. "
            "Return the data strictly in CSV format, where each row is on a new line "
            "and columns are separated by commas. "
            "Escape or wrap cells containing commas in double quotes. "
            "Return ONLY the CSV text with no markdown block formatting, markdown code ticks, "
            "or introductory text."
        )

        try:
            response = self.ocr.client.models.generate_content(
                model=self.ocr.model_name,
                contents=[
                    types.Part.from_bytes(
                        data=image_bytes,
                        mime_type="image/png",
                    ),
                    prompt,
                ],
            )
            raw_csv = response.text.replace("\r\n", "\n").strip()
            
            # Clean potential markdown block wrappers if the model ignores instructions
            if raw_csv.startswith("```"):
                lines_clean = raw_csv.splitlines()
                if lines_clean[0].startswith("```"):
                    lines_clean = lines_clean[1:]
                if lines_clean and lines_clean[-1].startswith("```"):
                    lines_clean = lines_clean[:-1]
                raw_csv = "\n".join(lines_clean).strip()

            if not raw_csv:
                return None

            import csv
            reader = csv.reader(io.StringIO(raw_csv))
            cell_rows = [row for row in reader if row]

            if not cell_rows:
                return None

            columns_count = max(len(row) for row in cell_rows)
            # Uniformize row column lengths
            for row in cell_rows:
                while len(row) < columns_count:
                    row.append("")

            column_widths = [float(width / columns_count)] * columns_count
            row_heights = [float(height / len(cell_rows))] * len(cell_rows)

            return TableData(
                rows=cell_rows,
                column_widths=column_widths,
                row_heights=row_heights,
            )

        except Exception as exc:
            logger.exception("Gemini Table OCR failed")
            return None


class DocxService:
    """Generate visually positioned DOCX pages from DocumentModel.

    Images use DrawingML anchors relative to the page.

    Text uses VML textboxes because WordprocessingML textboxes can contain
    ordinary WordprocessingML paragraphs/runs while retaining absolute
    shape geometry.

    Tables use the OOXML tblpPr floating-table mechanism.
    """

    EMU_PER_POINT = 12700
    TWIPS_PER_POINT = 20

    @classmethod
    def _emu(cls, points: float) -> int:
        return max(1, round(points * cls.EMU_PER_POINT))

    @classmethod
    def _twips(cls, points: float) -> int:
        return max(1, round(points * cls.TWIPS_PER_POINT))

    @staticmethod
    def _set_run_font(run: Any, font: FontSettings) -> None:
        run.font.name = font.family
        run.font.size = Pt(font.size_pt)
        run.font.bold = font.bold
        run.font.italic = font.italic
        run.font.underline = font.underline

        # Ensure the East Asia font mapping follows the selected family.
        rpr = run._r.get_or_add_rPr()
        rfonts = rpr.rFonts
        if rfonts is not None:
            rfonts.set(qn("w:ascii"), font.family)
            rfonts.set(qn("w:hAnsi"), font.family)
            rfonts.set(qn("w:eastAsia"), font.family)

    @classmethod
    def _add_anchor_to_run(
        cls,
        run: Any,
        x_pt: float,
        y_pt: float,
        width_pt: float,
        height_pt: float,
    ) -> bool:
        inline = run._r.xpath("./wp:drawing/wp:inline")
        if not inline:
            return False

        inline_element = inline[0]

        anchor = OxmlElement("wp:anchor")
        anchor.set("distT", "0")
        anchor.set("distB", "0")
        anchor.set("distL", "0")
        anchor.set("distR", "0")
        anchor.set("simplePos", "0")
        anchor.set("relativeHeight", "251658240")
        anchor.set("behindDoc", "0")
        anchor.set("locked", "0")
        anchor.set("layoutInCell", "1")
        anchor.set("allowOverlap", "1")

        simple_pos = OxmlElement("wp:simplePos")
        simple_pos.set("x", "0")
        simple_pos.set("y", "0")
        anchor.append(simple_pos)

        position_h = OxmlElement("wp:positionH")
        position_h.set("relativeFrom", "page")
        offset_h = OxmlElement("wp:posOffset")
        offset_h.text = str(cls._emu(x_pt))
        position_h.append(offset_h)
        anchor.append(position_h)

        position_v = OxmlElement("wp:positionV")
        position_v.set("relativeFrom", "page")
        offset_v = OxmlElement("wp:posOffset")
        offset_v.text = str(cls._emu(y_pt))
        position_v.append(offset_v)
        anchor.append(position_v)

        extent = OxmlElement("wp:extent")
        extent.set("cx", str(cls._emu(width_pt)))
        extent.set("cy", str(cls._emu(height_pt)))
        anchor.append(extent)

        effect = OxmlElement("wp:effectExtent")
        effect.set("l", "0")
        effect.set("t", "0")
        effect.set("r", "0")
        effect.set("b", "0")
        anchor.append(effect)

        wrap_none = OxmlElement("wp:wrapNone")
        anchor.append(wrap_none)

        for child in list(inline_element):
            if child.tag.endswith("extent"):
                continue
            anchor.append(child)

        inline_element.getparent().replace(inline_element, anchor)
        return True

    @classmethod
    def _add_positioned_image(
        cls,
        document: DocumentObject,
        image_bytes: bytes,
        x_pt: float,
        y_pt: float,
        width_pt: float,
        height_pt: float,
    ) -> None:
        if not image_bytes:
            return

        paragraph = document.add_paragraph()
        paragraph.paragraph_format.space_before = Pt(0)
        paragraph.paragraph_format.space_after = Pt(0)

        run = paragraph.add_run()
        try:
            stream = io.BytesIO(image_bytes)
            run.add_picture(
                stream,
                width=Pt(max(1.0, width_pt)),
                height=Pt(max(1.0, height_pt)),
            )
        except Exception:
            logger.exception("Failed to add picture to run")
            return

        success = cls._add_anchor_to_run(
            run,
            x_pt,
            y_pt,
            width_pt,
            height_pt,
        )
        if not success:
            logger.warning("Skipped image anchoring because DrawingML inline was missing.")

    @classmethod
    def _add_textbox(
        cls,
        document: DocumentObject,
        text: str,
        x_pt: float,
        y_pt: float,
        width_pt: float,
        height_pt: float,
        font: FontSettings,
    ) -> None:
        paragraph = document.add_paragraph()
        paragraph.paragraph_format.space_before = Pt(0)
        paragraph.paragraph_format.space_after = Pt(0)

        run = paragraph.add_run()

        # ------------------------------------------------------------
        # VML textbox
        # ------------------------------------------------------------

        pict = OxmlElement("w:pict")

        shape = OxmlElement("v:shape")

        shape.set(
            "style",
            (
                "position:absolute;"
                f"left:{x_pt:.2f}pt;"
                f"top:{y_pt:.2f}pt;"
                f"width:{width_pt:.2f}pt;"
                f"height:{height_pt:.2f}pt;"
                "z-index:1;"
                "mso-position-horizontal:absolute;"
                "mso-position-horizontal-relative:page;"
                "mso-position-vertical:absolute;"
                "mso-position-vertical-relative:page;"
            ),
        )

        shape.set("type", "#_x0000_t202")
        shape.set("stroked", "f")
        shape.set("filled", "f")

        # ------------------------------------------------------------
        # Textbox container
        # ------------------------------------------------------------

        textbox = OxmlElement("v:textbox")

        textbox.set(
            "style",
            "mso-fit-shape-to-text:false;",
        )

        content = OxmlElement("w:txbxContent")
        inner_p = OxmlElement("w:p")

        # ------------------------------------------------------------
        # Paragraph formatting
        # ------------------------------------------------------------

        ppr = OxmlElement("w:pPr")

        spacing = OxmlElement("w:spacing")
        spacing.set("before", "0")
        spacing.set("after", "0")
        spacing.set("line", "240")

        ppr.append(spacing)
        inner_p.append(ppr)

        # ------------------------------------------------------------
        # Text
        # ------------------------------------------------------------

        lines = text.splitlines() or [""]

        for index, line in enumerate(lines):
            inner_r = OxmlElement("w:r")
            rpr = OxmlElement("w:rPr")

            rfonts = OxmlElement("w:rFonts")
            rfonts.set(qn("w:ascii"), font.family)
            rfonts.set(qn("w:hAnsi"), font.family)
            rfonts.set(qn("w:eastAsia"), font.family)
            rpr.append(rfonts)

            size = OxmlElement("w:sz")
            size.set(
                "val",
                str(max(1, round(8))),
            )
            rpr.append(size)

            if font.bold:
                rpr.append(OxmlElement("w:b"))

            if font.italic:
                rpr.append(OxmlElement("w:i"))

            if font.underline:
                underline = OxmlElement("w:u")
                underline.set("val", "single")
                rpr.append(underline)

            inner_r.append(rpr)

            text_node = OxmlElement("w:t")
            text_node.text = line if line else " "

            if line.startswith(" ") or line.endswith(" "):
                text_node.set(
                    "{http://www.w3.org/XML/1998/namespace}space",
                    "preserve",
                )

            inner_r.append(text_node)

            if index < len(lines) - 1:
                br = OxmlElement("w:br")
                inner_r.append(br)

            inner_p.append(inner_r)

        content.append(inner_p)

        # ------------------------------------------------------------
        # Assemble VML textbox
        # ------------------------------------------------------------

        textbox.append(content)
        shape.append(textbox)

        wrap = OxmlElement("w10:wrap")
        wrap.set("type", "none")
        wrap.set("anchorx", "page")
        wrap.set("anchory", "page")

        shape.append(wrap)

        pict.append(shape)
        run._r.append(pict)


    @classmethod
    def _add_floating_table(
        cls,
        document: DocumentObject,
        table_data: TableData,
        x_pt: float,
        y_pt: float,
        width_pt: float,
        height_pt: float,
        font: FontSettings,
    ) -> None:
        rows = table_data.rows
        if not rows:
            return

        columns = max(len(row) for row in rows)
        if columns == 0:
            return

        table = document.add_table(
            rows=len(rows),
            cols=columns,
        )
        table.style = "Table Grid"
        table.autofit = False

        total_source_width = sum(table_data.column_widths)
        if total_source_width <= 0:
            column_widths = [width_pt / columns] * columns
        else:
            column_widths = [
                width_pt * source / total_source_width
                for source in table_data.column_widths
            ]

            if len(column_widths) < columns:
                column_widths.extend(
                    [width_pt / columns]
                    * (columns - len(column_widths))
                )

        total_source_height = sum(table_data.row_heights)
        for row_index, row in enumerate(rows):
            for column_index in range(columns):
                cell = table.cell(row_index, column_index)
                cell.width = Pt(column_widths[column_index])

                paragraph = cell.paragraphs[0]
                paragraph.alignment = WD_ALIGN_PARAGRAPH.LEFT
                paragraph.paragraph_format.space_before = Pt(0)
                paragraph.paragraph_format.space_after = Pt(0)

                value = (
                    row[column_index]
                    if column_index < len(row)
                    else ""
                )

                run = paragraph.add_run(value)
                cls._set_run_font(run, font)

            if row_index < len(table_data.row_heights):
                row_height = table_data.row_heights[row_index]
                rh_pt = (
                    height_pt * row_height / total_source_height
                    if total_source_height > 0
                    else height_pt / len(rows)
                )
                tr_pr = table.rows[row_index]._tr.get_or_add_trPr()
                tr_height = OxmlElement("w:trHeight")
                tr_height.set(
                    qn("w:val"),
                    str(cls._twips(rh_pt)),
                )
                tr_height.set(qn("w:hRule"), "atLeast")
                tr_pr.append(tr_height)

        tbl = table._tbl
        tbl_pr = tbl.tblPr

        tbl_layout = OxmlElement("w:tblLayout")
        tbl_layout.set(qn("w:type"), "fixed")
        tbl_pr.append(tbl_layout)

        tbl_position = OxmlElement("w:tblpPr")
        tbl_position.set(qn("w:leftFromText"), "0")
        tbl_position.set(qn("w:rightFromText"), "0")
        tbl_position.set(qn("w:topFromText"), "0")
        tbl_position.set(qn("w:bottomFromText"), "0")
        tbl_position.set(qn("w:vertAnchor"), "page")
        tbl_position.set(qn("w:horzAnchor"), "page")
        tbl_position.set(qn("w:tblpX"), str(cls._twips(x_pt)))
        tbl_position.set(qn("w:tblpY"), str(cls._twips(y_pt)))
        tbl_pr.append(tbl_position)

        tbl_width = OxmlElement("w:tblW")
        tbl_width.set(qn("w:type"), "dxa")
        tbl_width.set(qn("w:w"), str(cls._twips(width_pt)))
        tbl_pr.append(tbl_width)

    @classmethod
    def generate(
        cls,
        model: DocumentModel,
        output_path: str,
    ) -> None:
        if not model.pages:
            raise ValueError("There are no pages to generate.")

        output = Path(output_path)
        output.parent.mkdir(parents=True, exist_ok=True)

        document = Document()

        normal_style = document.styles["Normal"]
        normal_style.font.name = model.font.family
        normal_style.font.size = Pt(model.font.size_pt)

        for index, page in enumerate(model.pages):
            if index == 0:
                section = document.sections[0]
            else:
                section = document.add_section(WD_SECTION.NEW_PAGE)

            section.left_margin = Pt(0)
            section.right_margin = Pt(0)
            section.top_margin = Pt(0)
            section.bottom_margin = Pt(0)
            section.header_distance = Pt(0)
            section.footer_distance = Pt(0)

            section.page_width = Pt(page.width_pt)
            section.page_height = Pt(page.height_pt)

            # Anchor paragraphs and floating tables to this section/page.
            # An empty page therefore remains a real page in the document.
            document.add_paragraph()

            for element in page.elements:
                x_pt = element.x * page.width_pt
                y_pt = element.y * page.height_pt
                width_pt = element.width * page.width_pt
                height_pt = element.height * page.height_pt

                if element.element_type == ElementType.TEXT:
                    if not element.text.strip():
                        continue

                    cls._add_textbox(
                        document,
                        element.text,
                        x_pt,
                        y_pt,
                        width_pt,
                        height_pt,
                        model.font,
                    )

                elif element.element_type == ElementType.IMAGE:
                    if not element.image_png_b64:
                        continue

                    image_bytes = DocumentModel.b64_to_image(
                        element.image_png_b64
                    )

                    cls._add_positioned_image(
                        document,
                        image_bytes,
                        x_pt,
                        y_pt,
                        width_pt,
                        height_pt,
                    )

                elif element.element_type == ElementType.TABLE:
                    if (
                        element.table_fallback_image
                        and element.image_png_b64
                    ):
                        image_bytes = DocumentModel.b64_to_image(
                            element.image_png_b64
                        )
                        cls._add_positioned_image(
                            document,
                            image_bytes,
                            x_pt,
                            y_pt,
                            width_pt,
                            height_pt,
                        )
                    elif element.table:
                        cls._add_floating_table(
                            document,
                            element.table,
                            x_pt,
                            y_pt,
                            width_pt,
                            height_pt,
                            model.font,
                        )

        document.save(str(output))
        logger.info("DOCX generated: %s", output)