"""AutoPOE2 - Bộ Lọc Quang Học & Nhận Diện Nâng Cao (Vision Filters & Robust Matching).

Cung cấp các thuật toán xử lý ảnh nâng cao kháng nhiễu ánh sáng và độ phân giải:
- Dual-Pass CLAHE (Contrast Limited Adaptive Histogram Equalization) thuần NumPy.
- Sobel Gradient & Edge Magnitude cho trích xuất biên đặc trưng.
- Multi-Scale Normalized Cross-Correlation (NCC) Template Matching cho nút TRAVERSE.
- Nhận diện trạng thái Atlas UI và nút TRAVERSE với độ chính xác tuyệt đối, kháng bloom ánh sáng.
"""

from __future__ import annotations

import os
from typing import Optional, Tuple, Union

import numpy as np

try:
    from PIL import Image, ImageOps
except ImportError:
    Image = None  # type: ignore[assignment]

try:
    import winocr
except ImportError:
    winocr = None

from src.common.coordinate_transform import (
    CoordinateScaler,
    Resolution,
    ScalingMode,
    RES_1440P,
    get_traverse_button_pos,
)


def _load_pil(img_input: Union[str, Image.Image]) -> Optional[Image.Image]:
    """Chuyển đổi đường dẫn hoặc PIL Image thành RGB Image."""
    if img_input is None:
        return None
    if isinstance(img_input, str):
        if not os.path.exists(img_input):
            return None
        try:
            return Image.open(img_input).convert("RGB")
        except Exception:
            return None
    try:
        return img_input.convert("RGB")
    except Exception:
        return img_input


def apply_clahe_fast(
    img_gray: np.ndarray,
    grid_size: Tuple[int, int] = (8, 8),
    clip_limit: float = 2.5,
) -> np.ndarray:
    """Cân bằng biểu đồ tần suất thích ứng giới hạn tương phản (CLAHE) thuần NumPy.

    Triệt tiêu hiệu ứng bloom ánh sáng cam/vàng từ lò Map Device và tăng cường độ tương phản chữ.
    @param img_gray: Mảng 2D uint8 (H, W).
    @param grid_size: Số ô lưới chia ảnh (tiles_y, tiles_x).
    @param clip_limit: Hệ số giới hạn đỉnh histogram chống khuếch đại nhiễu.
    @return: Mảng 2D uint8 đã cân bằng sáng thích ứng.
    """
    h, w = img_gray.shape
    tiles_y, tiles_x = grid_size
    tile_h = max(1, h // tiles_y)
    tile_w = max(1, w // tiles_x)

    if tile_h < 4 or tile_w < 4:
        # Fallback cho ảnh quá nhỏ: Global Histogram Equalization có clip
        hist, _ = np.histogram(img_gray.flatten(), 256, [0, 256])
        clip_val = clip_limit * (h * w / 256.0)
        excess = np.sum(np.maximum(hist - clip_val, 0))
        clipped_hist = np.minimum(hist, clip_val) + (excess / 256.0)
        cdf = clipped_hist.cumsum()
        cdf_min = cdf.min()
        cdf_max = cdf.max()
        if cdf_max > cdf_min:
            cdf_norm = ((cdf - cdf_min) * 255.0 / (cdf_max - cdf_min)).astype(np.uint8)
            return cdf_norm[img_gray]
        return img_gray

    cdfs = np.zeros((tiles_y, tiles_x, 256), dtype=np.float32)
    clip_val = clip_limit * (tile_h * tile_w / 256.0)

    for ty in range(tiles_y):
        for tx in range(tiles_x):
            y_start = ty * tile_h
            y_end = (ty + 1) * tile_h if ty < tiles_y - 1 else h
            x_start = tx * tile_w
            x_end = (tx + 1) * tile_w if tx < tiles_x - 1 else w

            tile = img_gray[y_start:y_end, x_start:x_end]
            hist, _ = np.histogram(tile.flatten(), 256, [0, 256])

            excess = np.sum(np.maximum(hist - clip_val, 0))
            clipped = np.minimum(hist, clip_val) + (excess / 256.0)

            cdf = clipped.cumsum()
            cdf_min = cdf.min()
            cdf_max = cdf.max()
            if cdf_max > cdf_min:
                cdfs[ty, tx] = (cdf - cdf_min) * 255.0 / (cdf_max - cdf_min)
            else:
                cdfs[ty, tx] = np.linspace(0, 255, 256)

    # Vectorized bilinear interpolation
    y_coords = np.arange(h)
    x_coords = np.arange(w)

    ty_f = (y_coords - tile_h * 0.5) / float(tile_h)
    ty0 = np.floor(ty_f).astype(int)
    ty1 = ty0 + 1
    ay = (ty_f - ty0)[:, None]

    ty0_c = np.clip(ty0, 0, tiles_y - 1)
    ty1_c = np.clip(ty1, 0, tiles_y - 1)

    tx_f = (x_coords - tile_w * 0.5) / float(tile_w)
    tx0 = np.floor(tx_f).astype(int)
    tx1 = tx0 + 1
    ax = (tx_f - tx0)[None, :]

    tx0_c = np.clip(tx0, 0, tiles_x - 1)
    tx1_c = np.clip(tx1, 0, tiles_x - 1)

    TY0, TX0 = np.meshgrid(ty0_c, tx0_c, indexing="ij")
    TY1, TX1 = np.meshgrid(ty1_c, tx1_c, indexing="ij")

    c00 = cdfs[TY0, TX0, img_gray]
    c01 = cdfs[TY0, TX1, img_gray]
    c10 = cdfs[TY1, TX0, img_gray]
    c11 = cdfs[TY1, TX1, img_gray]

    top = (1.0 - ax) * c00 + ax * c01
    bot = (1.0 - ax) * c10 + ax * c11
    out = (1.0 - ay) * top + ay * bot

    return np.clip(out, 0, 255).astype(np.uint8)


def sobel_edge_magnitude(img_gray: np.ndarray) -> np.ndarray:
    """Tính toán biên độ gradient Sobel theo cả 2 trục x và y thuần NumPy."""
    f = img_gray.astype(np.float32)
    gx = (
        -f[:-2, :-2] + f[:-2, 2:]
        - 2.0 * f[1:-1, :-2] + 2.0 * f[1:-1, 2:]
        - f[2:, :-2] + f[2:, 2:]
    )
    gy = (
        -f[:-2, :-2] - 2.0 * f[:-2, 1:-1] - f[:-2, 2:]
        + f[2:, :-2] + 2.0 * f[2:, 1:-1] + f[2:, 2:]
    )
    mag = np.hypot(gx, gy)
    out = np.zeros_like(f)
    out[1:-1, 1:-1] = mag
    mx = out.max()
    return (out * 255.0 / mx).astype(np.uint8) if mx > 0 else out.astype(np.uint8)


def dual_pass_enhance_image(
    img: Union[Image.Image, np.ndarray],
    clip_limit: float = 2.5,
) -> Image.Image:
    """Áp dụng Dual-Pass CLAHE trên ảnh để tạo ảnh PIL tối ưu cho OCR."""
    if isinstance(img, Image.Image):
        gray = np.array(img.convert("L"))
    else:
        gray = img if img.ndim == 2 else np.dot(img[..., :3], [0.299, 0.587, 0.114]).astype(np.uint8)

    enhanced = apply_clahe_fast(gray, grid_size=(8, 8), clip_limit=clip_limit)
    return Image.fromarray(enhanced).convert("RGB")


def match_traverse_template_edge(
    full_img: Image.Image,
    template_img: Image.Image,
) -> Tuple[Optional[Tuple[int, int]], float]:
    """So khớp mẫu nút TRAVERSE dựa trên đặc trưng cạnh (Sobel Edge NCC).

    Kháng hoàn toàn biến động màu sắc, lửa lò Map Device và chói sáng bloom.
    @return: ((center_x, center_y), correlation_score)
    """
    w, h = full_img.size
    res = Resolution(w, h)
    scaler = CoordinateScaler(source_res=RES_1440P, target_res=res)

    tw = max(20, int(round(template_img.width * scaler.scale_x)))
    th = max(20, int(round(template_img.height * scaler.scale_y)))
    tpl_scaled = template_img.resize((tw, th), Image.Resampling.BICUBIC)

    # ROI vùng modal Atlas trung tâm (khoảng 32%..58% chiều rộng, 42%..65% chiều cao)
    # Tọa độ tham chiếu 1440p: x1=819.2, y1=604.8, x2=1484.8, y2=936.0
    mode = ScalingMode.CENTER_ANCHORED if res.is_windowed_aspect else ScalingMode.STRETCH
    rx1, ry1, rx2, ry2 = scaler.transform_bbox(819.2, 604.8, 1484.8, 936.0, mode=mode)
    max_traverse_x = int(w * 0.55)
    roi_x1 = max(0, min(w - 1, rx1))
    roi_y1 = max(0, min(h - 1, ry1))
    roi_x2 = max(roi_x1 + 1, min(max_traverse_x, rx2))
    roi_y2 = max(roi_y1 + 1, min(h, ry2))

    roi_img = full_img.crop((roi_x1, roi_y1, roi_x2, roi_y2)).convert("L")

    roi_edge = sobel_edge_magnitude(np.array(roi_img)).astype(np.float32)
    tpl_edge = sobel_edge_magnitude(np.array(tpl_scaled.convert("L"))).astype(np.float32)

    rh, rw = roi_edge.shape
    if rh < th or rw < tw:
        return None, 0.0

    tpl_mean = tpl_edge - tpl_edge.mean()
    tpl_norm = np.linalg.norm(tpl_mean)
    if tpl_norm < 1e-5:
        return None, 0.0

    best_corr = -1.0
    best_loc = (0, 0)

    # Coarse search stride 2
    for y in range(0, rh - th + 1, 2):
        for x in range(0, rw - tw + 1, 2):
            patch = roi_edge[y:y + th, x:x + tw]
            p_mean = patch - patch.mean()
            p_norm = np.linalg.norm(p_mean)
            if p_norm > 1e-5:
                c = float(np.sum(p_mean * tpl_mean) / (p_norm * tpl_norm))
                if c > best_corr:
                    best_corr = c
                    best_loc = (x, y)

    # Fine search stride 1 quanh best_loc
    bx, by = best_loc
    refine_x1 = max(0, bx - 3)
    refine_x2 = min(rw - tw, bx + 3)
    refine_y1 = max(0, by - 3)
    refine_y2 = min(rh - th, by + 3)

    for y in range(refine_y1, refine_y2 + 1):
        for x in range(refine_x1, refine_x2 + 1):
            patch = roi_edge[y:y + th, x:x + tw]
            p_mean = patch - patch.mean()
            p_norm = np.linalg.norm(p_mean)
            if p_norm > 1e-5:
                c = float(np.sum(p_mean * tpl_mean) / (p_norm * tpl_norm))
                if c > best_corr:
                    best_corr = c
                    best_loc = (x, y)

    tl_x = roi_x1 + best_loc[0]
    tl_y = roi_y1 + best_loc[1]

    # Tâm của từ TRAVERSE trong template chuẩn (offset 141.5px x 47.0px @ 2560x1440)
    text_cx = int(round(tl_x + 141.5 * scaler.scale_x))
    text_cy = int(round(tl_y + 47.0 * scaler.scale_y))

    if text_cx > max_traverse_x:
        return None, 0.0

    return (text_cx, text_cy), best_corr


def find_traverse_button_dual_pass(
    img_input: Union[str, Image.Image],
    template_path: Optional[str] = None,
    min_score: float = 0.55,
) -> Optional[Tuple[int, int]]:
    """Định vị nút TRAVERSE bằng kiến trúc Dual-Pass CLAHE + Canny Edge Template Matching.

    Pass 1: Thử nhận diện bằng WinOCR trực tiếp trên ảnh gốc.
    Pass 2: Khi Pass 1 thất bại do ánh sáng/bloom, kích hoạt CLAHE enhancement và
            Edge-based Normalized Cross-Correlation với template chuẩn.
    @param img_input: Đường dẫn tệp hoặc đối tượng PIL Image.
    @param template_path: Tùy chọn đường dẫn template; mặc định lấy `captures/traverse_button_crop.png`.
    @param min_score: Ngưỡng tương quan tối thiểu để xác thực (chống False-Positive).
    @return: (center_x, center_y) hoặc None nếu không phát hiện.
    """
    img = _load_pil(img_input)
    if img is None:
        return None

    w, h = img.size
    res = Resolution(w, h)
    scaler = CoordinateScaler(source_res=RES_1440P, target_res=res)
    mode = ScalingMode.CENTER_ANCHORED if res.is_windowed_aspect else ScalingMode.STRETCH
    rx1, ry1, rx2, ry2 = scaler.transform_bbox(819.2, 604.8, 1484.8, 936.0, mode=mode)
    
    # Bắt buộc: Kẹp chặt biên phải ROI tối đa 0.55 * w để loại trừ triệt để vùng túi đồ và tooltip Waystone
    max_traverse_x = int(w * 0.55)
    rx2_clamped = min(rx2, max_traverse_x)
    roi_box = (
        max(0, min(w - 1, rx1)),
        max(0, min(h - 1, ry1)),
        max(rx1 + 1, min(max_traverse_x, rx2_clamped)),
        max(ry1 + 1, min(h, ry2)),
    )

    # --- PASS 1: Direct WinOCR trên ROI nút bấm ---
    if winocr is not None:
        try:
            from src.common.vision_ocr import safe_recognize_pil_sync

            crop_roi = img.crop(roi_box)
            res = safe_recognize_pil_sync(crop_roi)
            for line in res.get("lines", []):
                t = line.get("text", "").upper()
                # Kháng bug chữ 'ENTER': Bỏ qua nếu là câu tooltip "allowing you to enter a map"
                if "ENTER" in t and "MAP" in t:
                    continue
                if "TRAVERSE" in t or "TRAVERS" in t:
                    words = [wd for wd in line.get("words", []) if "TRAV" in wd.get("text", "").upper()]
                    target = words if words else line.get("words", [])
                    if target:
                        min_x = min(wd["bounding_rect"]["x"] for wd in target)
                        max_x = max(wd["bounding_rect"]["x"] + wd["bounding_rect"]["width"] for wd in target)
                        min_y = min(wd["bounding_rect"]["y"] for wd in target)
                        max_y = max(wd["bounding_rect"]["y"] + wd["bounding_rect"]["height"] for wd in target)
                        cx = int(roi_box[0] + (min_x + max_x) / 2.0)
                        cy = int(roi_box[1] + (min_y + max_y) / 2.0)
                        if cx <= max_traverse_x:
                            return cx, cy

            # Thử Pass 1b: CLAHE Enhanced OCR
            enhanced_roi = dual_pass_enhance_image(crop_roi, clip_limit=3.0)
            res_enh = safe_recognize_pil_sync(enhanced_roi)
            for line in res_enh.get("lines", []):
                t = line.get("text", "").upper()
                if "ENTER" in t and "MAP" in t:
                    continue
                if "TRAVERSE" in t or "TRAVERS" in t:
                    words = [wd for wd in line.get("words", []) if "TRAV" in wd.get("text", "").upper()]
                    target = words if words else line.get("words", [])
                    if target:
                        min_x = min(wd["bounding_rect"]["x"] for wd in target)
                        max_x = max(wd["bounding_rect"]["x"] + wd["bounding_rect"]["width"] for wd in target)
                        min_y = min(wd["bounding_rect"]["y"] for wd in target)
                        max_y = max(wd["bounding_rect"]["y"] + wd["bounding_rect"]["height"] for wd in target)
                        cx = int(roi_box[0] + (min_x + max_x) / 2.0)
                        cy = int(roi_box[1] + (min_y + max_y) / 2.0)
                        if cx <= max_traverse_x:
                            return cx, cy
        except Exception:
            pass

    # --- PASS 2: Sobel Edge Multi-Scale Template Matching ---
    if template_path is None:
        default_tpl = os.path.join(os.path.dirname(__file__), "..", "..", "captures", "traverse_button_crop.png")
        template_path = os.path.abspath(default_tpl)

    if os.path.exists(template_path):
        try:
            tpl_img = Image.open(template_path)
            coords, score = match_traverse_template_edge(img, tpl_img)
            if coords and score >= min_score and coords[0] <= max_traverse_x:
                return coords
        except Exception:
            pass

    return None


def is_atlas_ui_open_robust(img_input: Union[str, Image.Image]) -> bool:
    """Kiểm tra trạng thái mở của giao diện Atlas/Map Device kháng nhiễu và kháng click mù.

    Yêu cầu:
    1. Phát hiện các token trên banner đỉnh ("SEARCH", "ENDGAME", "ACT 1-4", "ATLAS") HOẶC
    2. Định vị được nút TRAVERSE ở trung tâm qua Dual-Pass Locator.
    3. Chặn đứng False-Positives:
       - Từ chối nếu là hộp thoại NPC ("GOODBYE", "BUY OR SELL", "INTRODUCTION").
       - Từ chối nếu chỉ thấy chữ "MAP DEVICE" đơn lẻ trên sàn 3D Hideout mà không có banner/nút.
    """
    img = _load_pil(img_input)
    if img is None:
        return False

    w, h = img.size

    # 1. Kiểm tra nút TRAVERSE qua Dual-Pass Locator
    btn_pos = find_traverse_button_dual_pass(img)
    if btn_pos is not None:
        return True

    # 2. Kiểm tra banner đỉnh
    if winocr is not None:
        try:
            from src.common.vision_ocr import safe_recognize_pil_sync

            crop_banner = img.crop((0, 0, w, int(h * 0.16)))
            res_banner = safe_recognize_pil_sync(crop_banner)
            blob = " ".join(line.get("text", "") for line in res_banner.get("lines", [])).upper()

            if any(bad in blob for bad in ["GOODBYE", "BUY OR SELL", "INTRODUCTION"]):
                return False

            if any(k in blob for k in ["SEARCH", "ENDGAME", "INTERLUDE", "ACT 1", "ACT 2", "ACT 3", "ACT 4", "ATLAS", "WORLD"]):
                return True
        except Exception:
            pass

    return False
