#!/usr/bin/env python3
"""Disposable live bridge for the WC3 command-stream browser lab.

This is deliberately not a production service. It tails one native recorder
file, applies the reference relay, exposes the two packet classes on local
WebSockets, and injects browser input into the authoritative Xvfb game window.
Access it through an SSH tunnel; it binds only to loopback by default.
"""

from __future__ import annotations

import argparse
import asyncio
import json
import os
from pathlib import Path
import queue
import re
import signal
import stat
import subprocess
import sys
import threading
import time

import websockets
from Xlib import X, XK, display
from Xlib.ext import xtest


ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT))
from relay import CommandStreamRelay, RelayError  # noqa: E402


KEYMAP = {
    "ArrowUp": "Up",
    "ArrowDown": "Down",
    "ArrowLeft": "Left",
    "ArrowRight": "Right",
    " ": "space",
    "Escape": "Escape",
    "Enter": "Return",
    "Backspace": "BackSpace",
    "PageUp": "Prior",
    "PageDown": "Next",
    "+": "plus",
    "=": "equal",
    "-": "minus",
}


class InputInjector:
    def __init__(self, display_name: str):
        self.conn = display.Display(display_name)
        self.root = self.conn.screen().root
        self.x = 0
        self.y = 0
        self.width = 1024
        self.height = 768
        self.last_geometry_refresh = 0.0
        self.window_found = False

    def refresh_geometry(self, force: bool = False) -> None:
        now = time.monotonic()
        if not force and now - self.last_geometry_refresh < 1:
            return
        self.last_geometry_refresh = now

        def visit(window, parent_x=0, parent_y=0):
            try:
                name = window.get_wm_name() or ""
                geometry = window.get_geometry()
            except Exception:
                return None
            absolute_x = parent_x + geometry.x
            absolute_y = parent_y + geometry.y
            if "Warcraft III" in name:
                return (absolute_x, absolute_y,
                        geometry.width, geometry.height)
            try:
                children = window.query_tree().children
            except Exception:
                return None
            for child in children:
                found = visit(child, absolute_x, absolute_y)
                if found:
                    return found
            return None

        found = visit(self.root)
        if found:
            self.x, self.y, self.width, self.height = found
            self.window_found = True

    def geometry(self) -> dict[str, int | bool]:
        return {"windowFound": self.window_found, "originX": self.x,
                "originY": self.y, "width": self.width,
                "height": self.height}

    def move(self, nx: float, ny: float, force_refresh: bool = False) -> tuple[int, int]:
        self.refresh_geometry(force_refresh)
        x = self.x + max(0, min(
            self.width - 1, round(float(nx) * (self.width - 1))))
        y = self.y + max(0, min(
            self.height - 1, round(float(ny) * (self.height - 1))))
        xtest.fake_input(self.conn, X.MotionNotify, 0, X.CurrentTime,
                         self.root, x, y)
        self.conn.sync()
        return x, y

    def button(self, button: int, down: bool) -> None:
        if 1 <= button <= 7:
            xtest.fake_input(self.conn,
                             X.ButtonPress if down else X.ButtonRelease,
                             button)
            self.conn.sync()

    def wheel(self, delta: float) -> None:
        button = 4 if delta < 0 else 5
        xtest.fake_input(self.conn, X.ButtonPress, button)
        xtest.fake_input(self.conn, X.ButtonRelease, button)
        self.conn.sync()

    def key(self, name: str, down: bool) -> None:
        symbol = XK.string_to_keysym(KEYMAP.get(name, name))
        if not symbol and len(name) == 1:
            symbol = XK.string_to_keysym(name.lower())
        if not symbol:
            return
        code = self.conn.keysym_to_keycode(symbol)
        xtest.fake_input(self.conn, X.KeyPress if down else X.KeyRelease, code)
        self.conn.sync()


class LabBridge:
    ROLES = frozenset({"reliable", "frame", "control"})
    TOKEN = re.compile(r"^[A-Za-z0-9_-]{1,96}$")

    def __init__(self, capture: Path, display_name: str,
                 session_command: str | None = None):
        self.capture = capture
        self.display_name = display_name
        self.session_command = session_command
        self.reliable: queue.Queue[list[bytes]] = queue.Queue(maxsize=256)
        self.frames: queue.Queue[list[bytes]] = queue.Queue(maxsize=2)
        self.clients: dict[str, tuple[str, object]] = {}
        self.injector: InputInjector | None = None
        self.active_token: str | None = None
        self.started_token: str | None = None
        self.session_lock = asyncio.Lock()
        self.session_process: subprocess.Popen | None = None
        self.tail_process: subprocess.Popen | None = None
        self.pump_thread: threading.Thread | None = None
        self.pump_stop: threading.Event | None = None
        self.loop: asyncio.AbstractEventLoop | None = None
        self.disconnect_task: asyncio.Task | None = None
        self.session_monitor: asyncio.Task | None = None
        self.session_log = Path("/tmp/w3cs-session.log")

    @staticmethod
    def send_reliable(pending: queue.Queue[list[bytes]],
                      packets: list[bytes],
                      stop_event: threading.Event) -> bool:
        # Reliable resource/state records cannot be dropped. This pump reads a
        # recorder file, so bounded-queue backpressure stops only the reader;
        # it never blocks the game's recorder thread.
        while not stop_event.is_set():
            try:
                pending.put(packets, timeout=.1)
                return True
            except queue.Full:
                continue
        return False

    @staticmethod
    def send_frame(pending: queue.Queue[list[bytes]],
                   packets: list[bytes]) -> bool:
        while True:
            try:
                pending.put_nowait(packets)
                return True
            except queue.Full:
                try:
                    pending.get_nowait()
                except queue.Empty:
                    return False

    def pump(self, token: str, reliable: queue.Queue[list[bytes]],
             frames: queue.Queue[list[bytes]],
             stop_event: threading.Event) -> None:
        relay = CommandStreamRelay(
            lambda packets: self.send_reliable(
                reliable, packets, stop_event),
            lambda packets: self.send_frame(frames, packets))
        relay.stop_event = stop_event
        is_fifo = (self.capture.exists()
            and stat.S_ISFIFO(self.capture.stat().st_mode))
        command = (["cat", str(self.capture)] if is_fifo else
            ["tail", "--bytes=+1", "--follow=name", "--retry",
             "--sleep-interval=.005", str(self.capture)])
        tail = subprocess.Popen(
            command,
            stdout=subprocess.PIPE,
            stderr=subprocess.DEVNULL,
        )
        self.tail_process = tail
        if stop_event.is_set():
            tail.terminate()
            return
        try:
            assert tail.stdout is not None
            relay.pump(tail.stdout)
        except Exception as error:
            if not stop_event.is_set() and self.loop is not None:
                asyncio.run_coroutine_threadsafe(
                    self.fail_session(token, error), self.loop)
        finally:
            tail.terminate()
            if self.tail_process is tail:
                self.tail_process = None

    def stop_session_sync(self) -> None:
        if self.pump_stop is not None:
            self.pump_stop.set()
        tail = self.tail_process
        if tail is not None and tail.poll() is None:
            tail.terminate()
        process = self.session_process
        if process is not None and process.poll() is None:
            try:
                os.killpg(process.pid, signal.SIGTERM)
                process.wait(timeout=5)
            except (ProcessLookupError, subprocess.TimeoutExpired):
                try:
                    os.killpg(process.pid, signal.SIGKILL)
                except ProcessLookupError:
                    pass
        subprocess.run(["wineserver", "-k"], check=False,
                       stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        thread = self.pump_thread
        if thread is not None and thread.is_alive():
            thread.join(timeout=2)
        self.session_process = None
        self.tail_process = None
        self.pump_thread = None
        self.pump_stop = None

    async def monitor_session(self, token: str,
                              process: subprocess.Popen) -> None:
        code = await asyncio.to_thread(process.wait)
        if (token == self.active_token and token == self.started_token
                and self.pump_stop is not None
                and not self.pump_stop.is_set()):
            await self.fail_session(
                token, RuntimeError(
                    f"Warcraft exited with status {code}; see "
                    f"{self.session_log}"))

    async def stop_session(self) -> None:
        await asyncio.to_thread(self.stop_session_sync)
        monitor = self.session_monitor
        if monitor is not None and monitor is not asyncio.current_task():
            monitor.cancel()
        self.session_monitor = None
        self.started_token = None

    async def start_session(self, token: str) -> None:
        if self.started_token == token:
            return
        await self.stop_session()
        self.capture.parent.mkdir(parents=True, exist_ok=True)
        if not (self.capture.exists()
                and stat.S_ISFIFO(self.capture.stat().st_mode)):
            self.capture.unlink(missing_ok=True)
        reliable = self.reliable
        frames = self.frames
        stop_event = threading.Event()
        self.pump_stop = stop_event
        if self.session_command:
            session_log = self.session_log.open("wb")
            try:
                self.session_process = subprocess.Popen(
                    [self.session_command], start_new_session=True,
                    stdout=session_log, stderr=subprocess.STDOUT)
            finally:
                session_log.close()
            self.session_monitor = asyncio.create_task(
                self.monitor_session(token, self.session_process))
        self.pump_thread = threading.Thread(
            target=self.pump,
            args=(token, reliable, frames, stop_event), daemon=True)
        self.pump_thread.start()
        self.started_token = token
        print(f"session {token} started", flush=True)

    async def fail_session(self, token: str, error: Exception) -> None:
        async with self.session_lock:
            if token != self.active_token:
                return
            print(f"session {token} failed: {error}", flush=True)
            sockets = [entry[1] for entry in self.clients.values()
                       if entry[0] == token]
            for websocket in sockets:
                await websocket.close(code=1011,
                                      reason="native command stream failed")

    async def switch_session(self, token: str) -> None:
        if token == self.active_token:
            return
        old_sockets = [entry[1] for entry in self.clients.values()]
        self.clients.clear()
        self.active_token = token
        self.reliable = queue.Queue(maxsize=256)
        self.frames = queue.Queue(maxsize=2)
        if self.disconnect_task is not None:
            self.disconnect_task.cancel()
            self.disconnect_task = None
        for websocket in old_sockets:
            await websocket.close(code=1012, reason="new lab session started")
        await self.stop_session()

    async def expire_incomplete_session(self, token: str) -> None:
        try:
            await asyncio.sleep(.5)
            async with self.session_lock:
                if token != self.active_token:
                    return
                roles = {role for role, entry in self.clients.items()
                         if entry[0] == token}
                if roles == self.ROLES:
                    return
                sockets = [entry[1] for entry in self.clients.values()
                           if entry[0] == token]
                for websocket in sockets:
                    await websocket.close(
                        code=1012, reason="incomplete lab session")
                self.clients.clear()
                self.active_token = None
                await self.stop_session()
                print(f"session {token} stopped after disconnect", flush=True)
        except asyncio.CancelledError:
            pass

    async def sender(self, websocket, pending: queue.Queue[list[bytes]]) -> None:
        while True:
            try:
                packets = await asyncio.to_thread(pending.get, True, .5)
            except queue.Empty:
                continue
            for packet in packets:
                await websocket.send(packet)

    async def control(self, websocket) -> None:
        while self.injector is None:
            try:
                self.injector = InputInjector(self.display_name)
                self.injector.refresh_geometry()
            except Exception:
                await asyncio.sleep(0.25)
        async for raw in websocket:
            if not isinstance(raw, str):
                continue
            try:
                event = json.loads(raw)
                kind = event.get("t")
                applied = {"t": "inputAck", "kind": kind}
                if kind == "move":
                    applied["x"], applied["y"] = self.injector.move(
                        event["x"], event["y"])
                elif kind == "button":
                    applied["x"], applied["y"] = self.injector.move(
                        event["x"], event["y"], force_refresh=True)
                    self.injector.button(int(event.get("button", 1)),
                                         bool(event.get("down")))
                elif kind == "wheel":
                    applied["x"], applied["y"] = self.injector.move(
                        event["x"], event["y"], force_refresh=True)
                    self.injector.wheel(float(event.get("delta", 0)))
                elif kind == "key":
                    self.injector.key(str(event.get("key", "")),
                                      bool(event.get("down")))
                elif kind == "refreshGeometry":
                    self.injector.refresh_geometry(force=True)
                else:
                    continue
                applied.update(self.injector.geometry())
                await websocket.send(json.dumps(applied,
                                                separators=(",", ":")))
            except (KeyError, TypeError, ValueError):
                continue

    async def handler(self, websocket, path: str) -> None:
        parts = path.strip("/").split("/")
        if len(parts) == 1:
            role, token = parts[0], "legacy"
        elif len(parts) == 2:
            role, token = parts
        else:
            role, token = "", ""
        if role not in self.ROLES or not self.TOKEN.fullmatch(token):
            await websocket.close(code=1008, reason="unknown lab channel")
            return
        async with self.session_lock:
            await self.switch_session(token)
            previous = self.clients.get(role)
            if previous is not None:
                await previous[1].close(code=1012,
                                        reason="channel replaced")
            self.clients[role] = (token, websocket)
            roles = {name for name, entry in self.clients.items()
                     if entry[0] == token}
            if roles == self.ROLES:
                if self.disconnect_task is not None:
                    self.disconnect_task.cancel()
                    self.disconnect_task = None
                await self.start_session(token)
            pending = self.reliable if role == "reliable" else self.frames
        try:
            try:
                if role == "reliable":
                    await self.sender(websocket, pending)
                elif role == "frame":
                    await self.sender(websocket, pending)
                else:
                    await self.control(websocket)
            except websockets.exceptions.ConnectionClosed:
                pass
        finally:
            async with self.session_lock:
                current = self.clients.get(role)
                if current == (token, websocket):
                    del self.clients[role]
                    if self.disconnect_task is None:
                        self.disconnect_task = asyncio.create_task(
                            self.expire_incomplete_session(token))

    async def shutdown(self) -> None:
        async with self.session_lock:
            sockets = [entry[1] for entry in self.clients.values()]
            self.clients.clear()
            for websocket in sockets:
                await websocket.close(code=1001, reason="lab stopped")
            await self.stop_session()


async def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--capture", type=Path, default=Path("/tmp/w3cs.bin"))
    parser.add_argument("--display", default=":11")
    parser.add_argument("--host", default="127.0.0.1")
    parser.add_argument("--port", type=int, default=8145)
    parser.add_argument("--session-command")
    args = parser.parse_args()

    bridge = LabBridge(args.capture, args.display, args.session_command)
    bridge.loop = asyncio.get_running_loop()
    stopped = asyncio.Event()
    for name in (signal.SIGINT, signal.SIGTERM):
        try:
            bridge.loop.add_signal_handler(name, stopped.set)
        except NotImplementedError:
            pass
    try:
        async with websockets.serve(bridge.handler, args.host, args.port,
                                    max_size=64 * 1024 * 1024,
                                    ping_interval=10, ping_timeout=20):
            print(f"w3cs lab bridge listening on {args.host}:{args.port}",
                  flush=True)
            await stopped.wait()
    finally:
        await bridge.shutdown()


if __name__ == "__main__":
    asyncio.run(main())
