import os
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(".").resolve()))

from src.assistant_tool.client_launcher import Poe2ClientLauncher
from src.common.win32_window import find_poe2_window, ensure_window_focus, is_poe2_window_focused, get_client_rect
import scripts.reconnect_and_enter_hideout as reconnect_module
from src.assistant_tool.screen_capturer import ScreenCapturer
import src.common.win32_input as win32_input
import src.common.vision_ocr as vision_ocr
from PIL import Image

print("=== TEST LIVE WASD WALK & FOCUS ===")

hwnd = find_poe2_window()
if not hwnd:
    print("Launching game client...")
    Poe2ClientLauncher.launch_client(wait_timeout_sec=90)
    time.sleep(5)
    hwnd = find_poe2_window()

print(f"HWND: {hex(hwnd) if hwnd else None}")
if not hwnd:
    print("Failed to get HWND!")
    sys.exit(1)

# Ensure hideout
print("Verifying Hideout status...")
capturer = ScreenCapturer(output_dir="captures/live_test")
cap1 = capturer.capture_screenshot(event_name="PRE_WALK_CHECK", capture_ram=False, auto_focus=False, background_mode=True)
img1 = Image.open(cap1)
res1 = vision_ocr.safe_recognize_pil_sync(img1)
txt1 = " ".join(l.get("text", "") for l in res1.get("lines", [])).upper()
print(f"Tokens found: {txt1[:80]}...")

is_in_hideout = "SHIELD" in txt1 or "WAYPOINT" in txt1 or "MAP DEVICE" in txt1 or "WARD" in txt1

if not is_in_hideout:
    print("Not in hideout! Running reconnect...")
    rc = reconnect_module.main(hwnd=hwnd)
    print("Reconnect result:", rc.get("success"))
    time.sleep(3)
else:
    print("Already in hideout!")

# Focus window
print("Ensuring window focus...")
f_ok = ensure_window_focus(hwnd)
print("ensure_window_focus:", f_ok)
is_foc = is_poe2_window_focused(hwnd)
print("is_poe2_window_focused:", is_foc)

# Chụp ảnh trước khi bước
cap_before = capturer.capture_screenshot(event_name="BEFORE_WALK", capture_ram=False, auto_focus=False, background_mode=True)
img_b = Image.open(cap_before)

print("Sending WASD (Up/North: phím W) trong 1.5s...")
# Đi thẳng lên phía Bắc (screen_y > 0.38 -> phím W)
# dx=1.0, dy=1.0 -> screen_x=0, screen_y=1.414 -> Phím W
walk_ok = win32_input.send_wasd_direction(1.0, 1.0, dwell_ms=1.5, hwnd=hwnd)
print("send_wasd_direction result:", walk_ok)

time.sleep(0.5)

# Chụp ảnh sau khi bước
cap_after = capturer.capture_screenshot(event_name="AFTER_WALK", capture_ram=False, auto_focus=False, background_mode=True)
img_a = Image.open(cap_after)

# So sánh pixel diff
diffs = []
w, h = img_b.size
for x in range(w // 4, 3 * w // 4, 10):
    for y in range(h // 4, 3 * h // 4, 10):
        p1 = img_b.getpixel((x, y))[:3]
        p2 = img_a.getpixel((x, y))[:3]
        d = sum(abs(a - b) for a, b in zip(p1, p2)) / 3.0
        diffs.append(d)

avg_diff = sum(diffs) / len(diffs) if diffs else 0
max_diff = max(diffs) if diffs else 0
big_diffs = sum(1 for d in diffs if d > 20)
print(f"PIXEL DIFF: Avg={avg_diff:.2f}, Max={max_diff:.2f}, BigDiffCount={big_diffs}/{len(diffs)}")

# OCR so sánh nhãn Waypoint hoặc Map Device
res_b = vision_ocr.safe_recognize_pil_sync(img_b)
res_a = vision_ocr.safe_recognize_pil_sync(img_a)

def get_label_pos(ocr_res, label):
    for l in ocr_res.get("lines", []):
        if label in l.get("text", "").upper():
            words = l.get("words", [])
            if words:
                r = words[0].get("bounding_rect", {})
                return r.get("x"), r.get("y")
    return None, None

wp_b = get_label_pos(res_b, "WAYPOINT")
wp_a = get_label_pos(res_a, "WAYPOINT")
md_b = get_label_pos(res_b, "MAP DEVICE")
md_a = get_label_pos(res_a, "MAP DEVICE")

print(f"WAYPOINT pos: Before={wp_b} -> After={wp_a}")
print(f"MAP DEVICE pos: Before={md_b} -> After={md_a}")
