"""Image Preprocessing Engine for Multimodal Document OCR & Vision.

Provides pure Python/Pillow/Numpy algorithms for document deskewing,
contrast enhancement, salt-and-pepper noise reduction, adaptive binarization,
and PDF page rasterization without external C/OpenCV dependencies.
"""

from __future__ import annotations

import base64
import io
import logging
from typing import Any

import numpy as np
from PIL import Image, ImageEnhance, ImageFilter, ImageOps

logger = logging.getLogger("dscons.ocr.preprocessor")


class ImagePreprocessor:
    """Enterprise image preprocessor optimized for Vietnamese receipts, UNCs, and AEC documents."""

    @staticmethod
    def load_image(image_input: bytes | io.BytesIO | Image.Image) -> Image.Image:
        """Loads and normalizes an input into a RGB PIL Image."""
        if isinstance(image_input, Image.Image):
            return image_input.convert("RGB")
        if isinstance(image_input, (bytes, bytearray)):
            image_input = io.BytesIO(image_input)
        img = Image.open(image_input)
        return img.convert("RGB")

    @classmethod
    def estimate_skew_angle(
        cls,
        image: Image.Image,
        max_angle: float = 15.0,
        angle_step: float = 0.5,
    ) -> float:
        """Determines the skew angle of text lines using Horizontal Projection Profile variance."""
        try:
            # 1. Downscale to max dimension 800 for high-speed calculation (<50ms)
            w, h = image.size
            if max(w, h) > 800:
                scale = 800.0 / max(w, h)
                thumb = image.resize((int(w * scale), int(h * scale)), Image.Resampling.BILINEAR)
            else:
                thumb = image

            # 2. Convert to grayscale and binary edges
            gray = thumb.convert("L")
            # Apply edge filter to isolate text stroke lines
            edges = gray.filter(ImageFilter.FIND_EDGES)
            arr = np.array(edges, dtype=np.float32)

            best_score = -1.0
            best_angle = 0.0

            # Scan angles from -max_angle to +max_angle
            angles = np.arange(-max_angle, max_angle + angle_step, angle_step)
            for angle in angles:
                if abs(angle) < 1e-4:
                    rotated_arr = arr
                else:
                    rotated_img = Image.fromarray(arr.astype(np.uint8)).rotate(
                        float(angle),
                        resample=Image.Resampling.BILINEAR,
                        expand=False,
                        fillcolor=0,
                    )
                    rotated_arr = np.array(rotated_img, dtype=np.float32)

                # Horizontal projection profile: sum across columns (axis=1)
                profile = np.sum(rotated_arr, axis=1)
                # Text lines aligned with rows produce the highest profile variance
                score = float(np.var(profile))
                if score > best_score:
                    best_score = score
                    best_angle = float(angle)

            return best_angle
        except Exception as e:
            logger.warning("Deskew angle estimation error: %s", e)
            return 0.0

    @classmethod
    def deskew(
        cls,
        image: Image.Image,
        max_angle: float = 15.0,
        angle_step: float = 0.5,
    ) -> tuple[Image.Image, float]:
        """Auto-rotates skewed document image to be perfectly horizontal."""
        angle = cls.estimate_skew_angle(image, max_angle=max_angle, angle_step=angle_step)
        if abs(angle) < 0.3:
            return image, 0.0

        logger.info("Auto-deskewing document: rotating by %.1f degrees", -angle)
        # We rotate by -angle to straighten it back
        straightened = image.rotate(
            -angle,
            resample=Image.Resampling.BICUBIC,
            expand=True,
            fillcolor=(255, 255, 255),
        )
        return straightened, -angle

    @staticmethod
    def enhance_contrast(image: Image.Image, factor: float = 1.3) -> Image.Image:
        """Autocontrasts and stretches histogram to separate faint text from background."""
        # 1. Auto-contrast clipping 1% of outliers
        auto = ImageOps.autocontrast(image, cutoff=1)
        # 2. Additional contrast boost
        enhancer = ImageEnhance.Contrast(auto)
        return enhancer.enhance(factor)

    @staticmethod
    def denoise(image: Image.Image) -> Image.Image:
        """Removes salt-and-pepper noise and speckles using Median Filter."""
        return image.filter(ImageFilter.MedianFilter(size=3))

    @staticmethod
    def sharpen(image: Image.Image) -> Image.Image:
        """Sharpens character strokes for sharper OCR bounding boxes."""
        return image.filter(ImageFilter.SHARPEN)

    @classmethod
    def adaptive_binarize(cls, image: Image.Image) -> Image.Image:
        """Applies Otsu's optimal global thresholding to create crisp 1-bit black & white image."""
        gray = image.convert("L")
        arr = np.array(gray, dtype=np.uint8)

        # Compute histogram
        hist, _ = np.histogram(arr, bins=256, range=(0, 256))
        total = arr.size
        current_max, threshold = 0.0, 128
        sum_total = np.dot(np.arange(256), hist)
        sum_back, weight_back = 0.0, 0

        for t in range(256):
            weight_back += hist[t]
            if weight_back == 0:
                continue
            weight_fore = total - weight_back
            if weight_fore == 0:
                break

            sum_back += t * hist[t]
            mean_back = sum_back / weight_back
            mean_fore = (sum_total - sum_back) / weight_fore

            # Between-class variance
            var_between = float(weight_back) * float(weight_fore) * ((mean_back - mean_fore) ** 2)
            if var_between > current_max:
                current_max = var_between
                threshold = t

        bin_arr = np.where(arr > threshold, 255, 0).astype(np.uint8)
        return Image.fromarray(bin_arr, mode="L")

    @classmethod
    def preprocess_document_image(
        cls,
        image_input: bytes | io.BytesIO | Image.Image,
        apply_deskew: bool = True,
        apply_binarize: bool = False,
    ) -> Image.Image:
        """Full preprocessing pipeline: Load -> Contrast Enhance -> Denoise -> Deskew -> Sharpen."""
        img = cls.load_image(image_input)
        img = cls.enhance_contrast(img)
        img = cls.denoise(img)
        if apply_deskew:
            img, _ = cls.deskew(img)
        img = cls.sharpen(img)
        if apply_binarize:
            img = cls.adaptive_binarize(img)
        return img

    @staticmethod
    def render_pdf_to_images(
        pdf_bytes: bytes,
        max_pages: int = 5,
        dpi: int = 200,
    ) -> list[Image.Image]:
        """Rasterizes PDF pages into high-resolution PIL images using PyMuPDF."""
        images: list[Image.Image] = []
        try:
            import pymupdf
            doc = pymupdf.open(stream=pdf_bytes, filetype="pdf")
            num_pages = min(len(doc), max_pages)
            # 72 is standard PDF point size; scale = dpi / 72
            zoom = dpi / 72.0
            matrix = pymupdf.Matrix(zoom, zoom)

            for pno in range(num_pages):
                page = doc[pno]
                pix = page.get_pixmap(matrix=matrix, alpha=False)
                img = Image.frombytes("RGB", [pix.width, pix.height], pix.samples)
                images.append(img)
            doc.close()
        except Exception as e:
            logger.error("Failed to rasterize PDF to images: %s", e)
        return images

    @staticmethod
    def image_to_bytes(
        image: Image.Image,
        format: str = "JPEG",
        quality: int = 85,
        max_dimension: int = 2048,
    ) -> bytes:
        """Converts and resizes image to compressed bytes for AI Vision transfer."""
        w, h = image.size
        if max(w, h) > max_dimension:
            scale = float(max_dimension) / float(max(w, h))
            image = image.resize((int(w * scale), int(h * scale)), Image.Resampling.LANCZOS)

        out_buf = io.BytesIO()
        if format.upper() in ("JPEG", "JPG"):
            if image.mode != "RGB":
                image = image.convert("RGB")
            image.save(out_buf, format="JPEG", quality=quality, optimize=True)
        else:
            image.save(out_buf, format=format)
        return out_buf.getvalue()

    @classmethod
    def image_to_data_uri(
        cls,
        image: Image.Image,
        format: str = "JPEG",
        quality: int = 85,
    ) -> str:
        """Encodes PIL Image into Base64 Data URI string."""
        raw_bytes = cls.image_to_bytes(image, format=format, quality=quality)
        b64_str = base64.b64encode(raw_bytes).decode("ascii")
        mime = "image/jpeg" if format.upper() in ("JPEG", "JPG") else f"image/{format.lower()}"
        return f"data:{mime};base64,{b64_str}"
