# -*- coding: utf-8 -*-
"""w7rtc8790_91 - Windows 7 / Python 3.8 simple remote desktop."""
from __future__ import annotations

import asyncio
import io
import json
import logging
import secrets
import socket
import sys
import threading
import time
from functools import partial
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Optional
from urllib.parse import parse_qs, quote, urlparse

import tkinter as tk
from tkinter import messagebox

import mss
from PIL import Image
import win32api
import win32con

try:
    import websockets
except Exception as exc:
    websockets = None
    WEBSOCKETS_IMPORT_ERROR = exc
else:
    WEBSOCKETS_IMPORT_ERROR = None

from yjm_net_upnp import UpnpPortMapper

APP_NAME = "w7rtc8790_91"
HTTP_PORT = 8790
WS_PORT = 8791
FPS = 10
JPEG_QUALITY = 65
MAX_WIDTH = 1280
ROOT = Path(__file__).resolve().parent
WEB_DIR = ROOT / "web"
TOKEN_FILE = ROOT / "w7rtc8790_91.token"
LOG_FILE = ROOT / "w7rtc8790_91.log"


def get_or_create_token():
    try:
        if TOKEN_FILE.exists():
            value = TOKEN_FILE.read_text(encoding="utf-8").strip()
            if len(value) >= 16:
                return value
        value = secrets.token_urlsafe(24)
        TOKEN_FILE.write_text(value, encoding="utf-8")
        return value
    except Exception:
        return secrets.token_urlsafe(24)


def copy_text_windows(text):
    root = tk.Tk()
    root.withdraw()
    try:
        root.clipboard_clear()
        root.clipboard_append(text)
        root.update()
    finally:
        root.destroy()


class DesktopTarget(object):
    def __init__(self):
        self.rect = self._read_rect()

    @staticmethod
    def _read_rect():
        with mss.mss() as sct:
            mon = sct.monitors[0]
            return {
                "left": int(mon["left"]),
                "top": int(mon["top"]),
                "width": int(mon["width"]),
                "height": int(mon["height"]),
            }

    def refresh(self):
        self.rect = self._read_rect()


class RemoteController(object):
    def __init__(self, target, logger):
        self.target = target
        self.logger = logger
        self.mouse_down = False

    def _point(self, msg):
        rect = self.target.rect
        x = max(0.0, min(1.0, float(msg.get("x", 0))))
        y = max(0.0, min(1.0, float(msg.get("y", 0))))
        return int(rect["left"] + x * rect["width"]), int(rect["top"] + y * rect["height"])

    @staticmethod
    def _button_flags(button):
        if str(button) in ("2", "right"):
            return win32con.MOUSEEVENTF_RIGHTDOWN, win32con.MOUSEEVENTF_RIGHTUP
        if str(button) in ("1", "middle"):
            return win32con.MOUSEEVENTF_MIDDLEDOWN, win32con.MOUSEEVENTF_MIDDLEUP
        return win32con.MOUSEEVENTF_LEFTDOWN, win32con.MOUSEEVENTF_LEFTUP

    def handle(self, text):
        try:
            msg = json.loads(text)
            if not isinstance(msg, dict) or msg.get("type") != "control":
                return
            action = str(msg.get("action", ""))
            if action in ("move", "drag_move", "click", "dblclick", "rightclick", "mouse_down", "mouse_up"):
                x, y = self._point(msg)
                win32api.SetCursorPos((x, y))
            if action in ("move", "drag_move"):
                return
            if action == "mouse_down":
                down, _up = self._button_flags(msg.get("button", 0))
                win32api.mouse_event(down, 0, 0, 0, 0)
                self.mouse_down = True
            elif action == "mouse_up":
                _down, up = self._button_flags(msg.get("button", 0))
                win32api.mouse_event(up, 0, 0, 0, 0)
                self.mouse_down = False
            elif action in ("click", "rightclick"):
                button = 2 if action == "rightclick" else msg.get("button", 0)
                down, up = self._button_flags(button)
                win32api.mouse_event(down, 0, 0, 0, 0)
                win32api.mouse_event(up, 0, 0, 0, 0)
            elif action == "dblclick":
                down, up = self._button_flags(msg.get("button", 0))
                for _ in range(2):
                    win32api.mouse_event(down, 0, 0, 0, 0)
                    win32api.mouse_event(up, 0, 0, 0, 0)
            elif action == "wheel":
                delta = int(msg.get("deltaY", msg.get("delta", 0)))
                win32api.mouse_event(win32con.MOUSEEVENTF_WHEEL, 0, 0, -120 if delta > 0 else 120, 0)
            elif action == "key":
                self._key(msg)
            elif action == "text":
                self._text(str(msg.get("text", "")))
        except Exception as exc:
            self.logger.debug("control ignored: %s", exc)

    @staticmethod
    def _key(msg):
        key = str(msg.get("key", ""))
        vk_map = {
            "Enter": win32con.VK_RETURN, "Backspace": win32con.VK_BACK,
            "Tab": win32con.VK_TAB, "Escape": win32con.VK_ESCAPE,
            "Delete": win32con.VK_DELETE, "ArrowLeft": win32con.VK_LEFT,
            "ArrowRight": win32con.VK_RIGHT, "ArrowUp": win32con.VK_UP,
            "ArrowDown": win32con.VK_DOWN, "Home": win32con.VK_HOME,
            "End": win32con.VK_END, "PageUp": win32con.VK_PRIOR,
            "PageDown": win32con.VK_NEXT, " ": win32con.VK_SPACE,
        }
        vk = vk_map.get(key)
        if vk is None and len(key) == 1:
            vk = win32api.VkKeyScan(key) & 0xFF
        if vk is None:
            return
        mods = msg.get("mods") or {}
        held = []
        for name, code in (("ctrl", win32con.VK_CONTROL), ("alt", win32con.VK_MENU), ("shift", win32con.VK_SHIFT)):
            if mods.get(name):
                win32api.keybd_event(code, 0, 0, 0)
                held.append(code)
        win32api.keybd_event(vk, 0, 0, 0)
        win32api.keybd_event(vk, 0, win32con.KEYEVENTF_KEYUP, 0)
        for code in reversed(held):
            win32api.keybd_event(code, 0, win32con.KEYEVENTF_KEYUP, 0)

    @staticmethod
    def _text(text):
        for ch in text[:500]:
            vk = win32api.VkKeyScan(ch)
            if vk == -1:
                continue
            code = vk & 0xFF
            shift = bool(vk & 0x0100)
            if shift:
                win32api.keybd_event(win32con.VK_SHIFT, 0, 0, 0)
            win32api.keybd_event(code, 0, 0, 0)
            win32api.keybd_event(code, 0, win32con.KEYEVENTF_KEYUP, 0)
            if shift:
                win32api.keybd_event(win32con.VK_SHIFT, 0, win32con.KEYEVENTF_KEYUP, 0)


class WebSocketFrameServer(object):
    def __init__(self, token, controller, logger):
        self.token = token
        self.controller = controller
        self.logger = logger
        self.clients = set()
        self.loop = None
        self.thread = None
        self.ready = threading.Event()
        self.running = False

    async def handler(self, websocket, path):
        got = parse_qs(urlparse(path).query).get("token", [""])[0]
        if not secrets.compare_digest(str(got), str(self.token)):
            await websocket.close(code=1008, reason="invalid token")
            return
        queue = asyncio.Queue(maxsize=1)
        self.clients.add(queue)

        async def sender():
            while True:
                await websocket.send(await queue.get())

        async def receiver():
            async for message in websocket:
                if isinstance(message, str):
                    self.controller.handle(message)

        send_task = asyncio.ensure_future(sender())
        recv_task = asyncio.ensure_future(receiver())
        try:
            _done, pending = await asyncio.wait([send_task, recv_task], return_when=asyncio.FIRST_COMPLETED)
            for task in pending:
                task.cancel()
        finally:
            self.clients.discard(queue)

    async def run(self):
        if websockets is None:
            raise RuntimeError("websockets import 실패: %s" % WEBSOCKETS_IMPORT_ERROR)
        async with websockets.serve(self.handler, "0.0.0.0", WS_PORT, max_size=None, compression=None, ping_interval=20):
            self.ready.set()
            await asyncio.Future()

    def start(self):
        self.running = True
        self.thread = threading.Thread(target=self._thread_main, daemon=True)
        self.thread.start()

    def _thread_main(self):
        self.loop = asyncio.new_event_loop()
        asyncio.set_event_loop(self.loop)
        try:
            self.loop.run_until_complete(self.run())
        except Exception as exc:
            self.logger.exception("WebSocket server failed: %s", exc)
            self.ready.set()

    def push(self, data):
        if self.loop and self.running:
            self.loop.call_soon_threadsafe(self._push_loop, data)

    def _push_loop(self, data):
        for queue in list(self.clients):
            try:
                if queue.full():
                    queue.get_nowait()
                queue.put_nowait(data)
            except Exception:
                self.clients.discard(queue)


class CaptureWorker(object):
    def __init__(self, target, ws_server, logger):
        self.target = target
        self.ws_server = ws_server
        self.logger = logger
        self.running = False
        self.thread = None

    def start(self):
        self.running = True
        self.thread = threading.Thread(target=self.run, daemon=True)
        self.thread.start()

    def run(self):
        with mss.mss() as sct:
            while self.running:
                started = time.time()
                try:
                    shot = sct.grab(self.target.rect)
                    image = Image.frombytes("RGB", shot.size, shot.bgra, "raw", "BGRX")
                    if image.width > MAX_WIDTH:
                        h = max(1, int(image.height * float(MAX_WIDTH) / float(image.width)))
                        resample = getattr(Image, "Resampling", Image).BILINEAR
                        image = image.resize((MAX_WIDTH, h), resample)
                    buf = io.BytesIO()
                    image.save(buf, "JPEG", quality=JPEG_QUALITY)
                    self.ws_server.push(buf.getvalue())
                except Exception as exc:
                    self.logger.debug("capture error: %s", exc)
                    time.sleep(0.5)
                time.sleep(max(0.001, (1.0 / FPS) - (time.time() - started)))


class HttpServer(object):
    def __init__(self):
        self.httpd = None
        self.thread = None

    def start(self):
        handler = partial(SimpleHTTPRequestHandler, directory=str(WEB_DIR))
        self.httpd = ThreadingHTTPServer(("0.0.0.0", HTTP_PORT), handler)
        self.thread = threading.Thread(target=self.httpd.serve_forever, daemon=True)
        self.thread.start()

    def stop(self):
        if self.httpd:
            self.httpd.shutdown()
            self.httpd.server_close()


class SimpleApp(object):
    def __init__(self, root):
        self.root = root
        self.root.title(APP_NAME)
        self.root.geometry("580x290")
        self.root.resizable(False, False)
        self.logger = logging.getLogger(APP_NAME)
        self.token = get_or_create_token()
        self.target = DesktopTarget()
        self.controller = RemoteController(self.target, self.logger)
        self.ws_server = WebSocketFrameServer(self.token, self.controller, self.logger)
        self.http_server = HttpServer()
        self.capture = CaptureWorker(self.target, self.ws_server, self.logger)
        self.upnp = UpnpPortMapper(self.logger)
        self.url = ""
        self.status = tk.StringVar(value="시작 중...")
        self.url_var = tk.StringVar(value="")
        self._build_ui()
        self.root.protocol("WM_DELETE_WINDOW", self.close)
        self.root.after(100, self.start_all)

    def _build_ui(self):
        tk.Label(self.root, text=APP_NAME, font=("Arial", 18, "bold")).pack(pady=(18, 8))
        tk.Label(self.root, textvariable=self.status, anchor="w", justify="left").pack(fill="x", padx=22, pady=5)
        entry = tk.Entry(self.root, textvariable=self.url_var, font=("Arial", 10))
        entry.pack(fill="x", padx=22, pady=8)
        row = tk.Frame(self.root)
        row.pack(pady=12)
        tk.Button(row, text="URL 복사", width=14, command=self.copy_url).pack(side="left", padx=5)
        tk.Button(row, text="UPnP 다시 열기", width=14, command=self.reopen_upnp).pack(side="left", padx=5)
        tk.Button(row, text="종료", width=14, command=self.close).pack(side="left", padx=5)
        tk.Label(self.root, text="HTTP 8790 / WebSocket 8791 · 전체 바탕화면 · 원격 마우스/키보드", fg="#555").pack(pady=6)

    def start_all(self):
        try:
            self.http_server.start()
            self.ws_server.start()
            self.ws_server.ready.wait(5.0)
            self.capture.start()
            self.reopen_upnp()
        except Exception as exc:
            self.status.set("실행 실패: %s" % exc)
            self.logger.exception("startup failed")

    def reopen_upnp(self):
        self.status.set("UPnP 포트 개방 중...")
        self.root.update_idletasks()
        lan_ip = UpnpPortMapper.get_lan_ip()
        result = self.upnp.open_ports(lan_ip, HTTP_PORT, WS_PORT, "https://api.ipify.org", APP_NAME)
        router_ip = result.router_external_ip
        public_ip = router_ip
        if router_ip and self.upnp.is_private_or_cgnat(router_ip):
            public_ip = result.echo_public_ip or router_ip
        elif not public_ip:
            public_ip = result.echo_public_ip
        if result.ok and public_ip:
            ws_url = "ws://%s:%d/?token=%s" % (public_ip, WS_PORT, quote(self.token))
            self.url = "http://%s:%d/viewer.html?ws=%s" % (public_ip, HTTP_PORT, quote(ws_url, safe=""))
            self.url_var.set(self.url)
            self.copy_url(silent=True)
            warning = ""
            if result.router_external_ip and result.echo_public_ip and result.router_external_ip != result.echo_public_ip:
                warning = "\n주의: 이중 NAT/CGNAT 가능성"
            self.status.set("실행 중 · UPnP 성공 · 외부 URL을 클립보드에 복사했습니다.%s" % warning)
        else:
            local_ws = "ws://%s:%d/?token=%s" % (lan_ip, WS_PORT, quote(self.token))
            self.url = "http://%s:%d/viewer.html?ws=%s" % (lan_ip, HTTP_PORT, quote(local_ws, safe=""))
            self.url_var.set(self.url)
            self.copy_url(silent=True)
            self.status.set("UPnP 실패 · LAN URL을 복사했습니다.\n%s" % result.message)

    def copy_url(self, silent=False):
        if not self.url:
            return
        self.root.clipboard_clear()
        self.root.clipboard_append(self.url)
        self.root.update()
        if not silent:
            self.status.set("URL을 클립보드에 다시 복사했습니다.")

    def close(self):
        try:
            self.capture.running = False
            self.ws_server.running = False
            self.http_server.stop()
            self.upnp.close_ports()
        finally:
            self.root.destroy()


def setup_logging():
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s %(levelname)s %(message)s",
        handlers=[logging.FileHandler(str(LOG_FILE), encoding="utf-8"), logging.StreamHandler(sys.stdout)],
    )


def main():
    setup_logging()
    root = tk.Tk()
    SimpleApp(root)
    root.mainloop()
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
