import sys
sys.path.insert(0, ".")
from PIL import Image
import src.common.vision_ocr as vision_ocr
from src.common.vision_ocr import (
    _load_image, CoordinateScaler, Resolution, RES_1440P,
    find_portal_label_by_ocr, find_color_clusters, is_loading_screen,
    safe_recognize_pil_sync, winocr
)

img = Image.open("captures/pipeline_run/20260919_171933_678_STAGE2_MAP_DEVICE_VIEW_APPROACH_3.png")
scale_x = 1197 / 2560.0
scale_y = 897 / 1440.0
fallback_pos = (1650, 140)

w, h = img.size
res = Resolution(w, h)
scaler = CoordinateScaler(source_res=RES_1440P, target_res=res)

sx = scale_x
sy = scale_y
default_x = int(round(fallback_pos[0] * scale_x))
default_y = int(round(fallback_pos[1] * scale_y))
print(f"default_x={default_x}, default_y={default_y}")

# 1. OCR Tìm nhãn 'MAP DEVICE'
bench_pts = []
if winocr is not None:
    try:
        regions = [
            (0, 0, int(w * 0.60), int(h * 0.70)),
            (int(w * 0.35), 0, int(w * 0.90), int(h * 0.70)),
            (0, 0, w, int(h * 0.85)),
        ]
        for idx, (rx1, ry1, rx2, ry2) in enumerate(regions):
            crop_img = img.crop((rx1, ry1, rx2, ry2))
            res = safe_recognize_pil_sync(crop_img)
            for line in res.get("lines", []):
                text = line.get("text", "").upper()
                clean = text.replace(" ", "").replace("'", "")
                if (
                    "DEVICE" in clean
                    or "MAPDEV" in clean
                    or ("MAPD" in clean and any(k in clean for k in ["EVI", "FC", "TNC", "RR"]))
                    or any(d in clean for d in ["UEVICE", "WEVICE", "EVICE"])
                    or ("MAP" in text and any(d in clean for d in ["DEV", "DEVI", "EVI"]))
                ):
                    print(f"Step 1 matched in region {idx}: line='{text}', clean='{clean}'")
                    words = line.get("words", [])
                    valid_words = [
                        w_info for w_info in words
                        if not ((w_info["bounding_rect"]["x"] + rx1) > 0.90 * w and (w_info["bounding_rect"]["y"] + ry1) < 0.18 * h)
                        and 0.02 * w <= (w_info["bounding_rect"]["x"] + rx1) <= 0.98 * w
                    ]
                    matched_words = [
                        w_info for w_info in valid_words
                        if any(k in w_info.get("text", "").upper() for k in ["MAP", "DEV", "DEVICE", "UEVICE", "WEVICE", "EVICE"])
                        or w_info.get("text", "").upper().replace(" ", "").replace("'", "") in ["D", "EVI", "FC", "TNC"]
                    ]
                    target_words = matched_words if matched_words else valid_words
                    if target_words:
                        min_x = min(w_info["bounding_rect"]["x"] for w_info in target_words) + rx1
                        max_x = max(w_info["bounding_rect"]["x"] + w_info["bounding_rect"]["width"] for w_info in target_words) + rx1
                        min_y = min(w_info["bounding_rect"]["y"] for w_info in target_words) + ry1
                        max_y = max(w_info["bounding_rect"]["y"] + w_info["bounding_rect"]["height"] for w_info in target_words) + ry1
                        cx = int((min_x + max_x) / 2)
                        cy = int((min_y + max_y) / 2)
                        print(f"Step 1 RETURN: cx={cx}, cy={cy}")

        # 1b
        p_labels = find_portal_label_by_ocr(img, scale_x=sx, scale_y=sy)
        print("p_labels:", p_labels)
        if p_labels:
            px = sum(p[0] for p in p_labels) // len(p_labels)
            py = sum(p[1] for p in p_labels) // len(p_labels)
            print(f"Step 1b RETURN: px={px}, py={py}")

        # 1c
        for line in res.get("lines", []):
            text = line.get("text", "").upper()
            if any(k in text for k in ["WAYPOINT", "WAYPOIN", "WAYPOI"]):
                print("1c Waypoint found:", text)
    except Exception as e:
        print("OCR error:", e)

# 2. Cyan
pedestal_roi = (
    max(0, int(default_x - 180 * sx)),
    max(0, int(default_y - 160 * sy)),
    min(w, int(default_x + 180 * sx)),
    min(h, int(default_y + 160 * sy)),
)
def is_cyan_gem(r: int, g: int, b: int) -> bool:
    return (g >= 110 and b >= 110 and g > r + 20 and b > r + 10) or \
           (g >= 135 and g > r * 1.35 and g > b * 1.1)
cyan_clusters = find_color_clusters(
    img, color_predicate=is_cyan_gem, roi=pedestal_roi,
    cluster_radius_x=30 * sx, cluster_radius_y=30 * sy, min_cluster_size=20,
)
print("cyan_clusters:", cyan_clusters)

# 3. Flame
def is_flame(r: int, g: int, b: int) -> bool:
    return r >= 190 and 90 <= g <= 220 and b <= 130 and (r > b + 50)
flame_roi = (
    max(0, int(default_x - 140 * sx)),
    max(0, int(default_y - 120 * sy)),
    min(w, int(default_x + 140 * sx)),
    min(h, int(default_y + 120 * sy)),
)
clusters = find_color_clusters(
    img, color_predicate=is_flame, roi=flame_roi,
    cluster_radius_x=25 * sx, cluster_radius_y=25 * sy, min_cluster_size=20,
)
print("flame clusters in roi:", clusters)
