"""
FreeExile WebSocket Gateway Bridge.
Translates W3C browser WebSocket frames (RFC 6455) to internal server actor packets,
supporting dual-mode framing (2-byte Opcode Protobuf and legacy JSON), live RTT telemetry,
authoritative combat actions, and zero-residual momentum verification.
"""

from __future__ import annotations

import asyncio
import json
import logging
import math
import struct
import sys
import threading
import time
from pathlib import Path
from typing import Any, Dict, Optional, Set, Tuple

# Ensure server package path resolution for proto imports
_SERVER_ROOT = Path(__file__).resolve().parent.parent
if str(_SERVER_ROOT) not in sys.path:
    sys.path.insert(0, str(_SERVER_ROOT))

try:
    from proto import auth_pb2, chat_pb2, combat_pb2, map_zone_pb2, network_pb2
except ImportError:
    network_pb2, combat_pb2, map_zone_pb2, chat_pb2, auth_pb2 = None, None, None, None, None  # type: ignore

try:
    from server.gateway.gateway_sync_helper import GatewaySyncHelper
except ImportError:
    from gateway.gateway_sync_helper import GatewaySyncHelper

try:
    import websockets
    from websockets.server import WebSocketServerProtocol
except ImportError:
    websockets = None  # type: ignore
    WebSocketServerProtocol = Any  # type: ignore

logger = logging.getLogger("WsGatewayBridge")

# Canonical 2-Byte Big-Endian Opcodes (Milestone M3 Specification & OpcodeRegistry.ts)
OPCODE_PACKET_ENVELOPE = 0x0001
OPCODE_PLAYER_MOVE_INPUT = 0x0010
OPCODE_ENTITY_STATE = 0x0011
OPCODE_ZONE_DATA = 0x0012
OPCODE_CAST_MARTIAL_SKILL = 0x0020
OPCODE_PHANTOM_EVASION = 0x0021
OPCODE_COMBAT_DAMAGE_EVENT = 0x0022
OPCODE_ENTER_ZONE_REQUEST = 0x0030
OPCODE_ZONE_PORTAL_DATA = 0x0031
OPCODE_CHAT_MESSAGE = 0x0040
OPCODE_AUTH_REQUEST = 0x0050


def encode_binary_frame(opcode: int, payload: bytes) -> bytes:
    """Encodes a binary frame with 2-byte big-endian uint16 opcode header."""
    return struct.pack(">H", opcode) + payload


def decode_binary_frame(frame: bytes) -> Tuple[int, bytes]:
    """Decodes a binary frame into (opcode, payload)."""
    if len(frame) < 2:
        raise ValueError("Binary frame too short for 2-byte opcode header")
    return struct.unpack(">H", frame[:2])[0], frame[2:]


class WsGatewayBridge:
    def __init__(self, host: str = "127.0.0.1", port: int = 8080):
        self.host = host
        self.port = port
        self.is_running = False
        self.server = None
        self.clients: Set[WebSocketServerProtocol] = set()
        self.loop: Optional[asyncio.AbstractEventLoop] = None
        self._thread: Optional[threading.Thread] = None
        self.player_states: Dict[int, Dict[str, Any]] = {}
        self.sync_helper = GatewaySyncHelper()
        try:
            from server.trade.two_phase_commit import InstantBuyoutEngine
            from server.world.agent_orb_service import AgentOrbService
        except ImportError:
            from trade.two_phase_commit import InstantBuyoutEngine
            from world.agent_orb_service import AgentOrbService
        self.buyout_engine = InstantBuyoutEngine(db_path="data/trade_ledger.db")
        self.orb_service = AgentOrbService(db_path="data/agent_delegations.db")

    def _get_player_state(self, entity_id: int) -> Dict[str, Any]:
        """Returns authoritative server state for an entity, creating if missing."""
        if entity_id not in self.player_states:
            self.player_states[entity_id] = {
                "x": 0.0, "y": 0.0, "vx": 0.0, "vy": 0.0, "speed": 6.0,
                "hp": 1000, "max_hp": 1000, "anim_state": 0,
                "last_time": time.time(), "last_seq": 0,
            }
        return self.player_states[entity_id]

    def verify_zero_residual_momentum(
        self, entity_id: int, input_x: float, input_y: float
    ) -> Tuple[bool, float, float]:
        """
        Zero-Residual Momentum verification:
        When input_x == 0 and input_y == 0, velocity instantly snaps to (0.0, 0.0) in frame 0.
        Returns (is_halted, vel_x, vel_y).
        """
        mag = math.hypot(input_x, input_y)
        is_halted = mag <= 0.05
        player = self._get_player_state(entity_id)
        if is_halted:
            player["vx"], player["vy"], player["anim_state"] = 0.0, 0.0, 0  # 0=Idle
            return True, 0.0, 0.0
        norm_x = input_x / mag if mag > 0 else 0.0
        norm_y = input_y / mag if mag > 0 else 0.0
        speed = player.get("speed", 6.0)
        player["vx"] = norm_x * speed
        player["vy"] = norm_y * speed
        player["anim_state"] = 1  # 1=Run
        return False, player["vx"], player["vy"]

    def _process_move_input(self, payload: bytes) -> Any:
        """Processes PlayerMoveInput and returns authoritative EntitySnapshot."""
        move_input = network_pb2.PlayerMoveInput.FromString(payload)
        entity_id = move_input.entity_id or 1001
        is_halted, vx, vy = self.verify_zero_residual_momentum(
            entity_id, move_input.dir_x, move_input.dir_y
        )
        player = self._get_player_state(entity_id)
        now = time.time()
        dt = min(0.1, max(0.001, now - player["last_time"]))
        player["last_time"] = now
        if not is_halted:
            player["x"] += vx * dt
            player["y"] += vy * dt
        player["last_seq"] = move_input.input_sequence
        return network_pb2.EntitySnapshot(
            entity_id=entity_id, pos_x=player["x"], pos_y=player["y"],
            velocity_x=vx, velocity_y=vy, current_hp=int(player["hp"]),
            max_hp=int(player["max_hp"]), animation_state=player["anim_state"],
            last_processed_input_seq=move_input.input_sequence,
        )

    async def _handle_binary_packet(
        self, websocket: WebSocketServerProtocol, opcode: int, payload: bytes
    ) -> None:
        """Dispatches binary protobuf frames by 2-byte opcode."""
        if network_pb2 is None:
            return

        if opcode == OPCODE_PLAYER_MOVE_INPUT:
            snap = self._process_move_input(payload)
            await websocket.send(encode_binary_frame(OPCODE_ENTITY_STATE, snap.SerializeToString()))
        elif opcode == OPCODE_PACKET_ENVELOPE:
            env = network_pb2.PacketEnvelope.FromString(payload)
            resp = network_pb2.PacketEnvelope(
                sequence_number=env.sequence_number, timestamp_ms=int(time.time() * 1000), nonce=env.nonce,
            )
            await websocket.send(encode_binary_frame(OPCODE_PACKET_ENVELOPE, resp.SerializeToString()))
        elif opcode == OPCODE_ENTITY_STATE:
            snap = network_pb2.EntitySnapshot.FromString(payload)
            player = self._get_player_state(snap.entity_id)
            player["x"], player["y"] = snap.pos_x, snap.pos_y
            player["vx"], player["vy"], player["anim_state"] = snap.velocity_x, snap.velocity_y, snap.animation_state
            await websocket.send(encode_binary_frame(OPCODE_ENTITY_STATE, snap.SerializeToString()))
        elif opcode == OPCODE_ZONE_DATA:
            zone = map_zone_pb2.ZoneInfoPayload(
                zone_id="zone_boundless_sanctuary", name="Doanh Trại Bến Lưu Đày",
                zone_type=map_zone_pb2.ZoneType.ZONE_TYPE_SANCTUARY, bounds_width=1600.0, bounds_height=1600.0,
            )
            await websocket.send(encode_binary_frame(OPCODE_ZONE_DATA, zone.SerializeToString()))
        elif opcode == OPCODE_CAST_MARTIAL_SKILL:
            skill_req = combat_pb2.CastMartialSkillRequest.FromString(payload)
            dmg_evt = combat_pb2.CombatDamageEvent(
                source_entity_id=skill_req.caster_entity_id or 1001, target_entity_id=9999,
                raw_damage=180, mitigated_damage=150, is_critical=False, target_evaded=False,
                element=combat_pb2.FiveElementsType.ELEMENT_PHYSICAL_KIM, recoil_damage_to_source=0,
            )
            await websocket.send(encode_binary_frame(OPCODE_COMBAT_DAMAGE_EVENT, dmg_evt.SerializeToString()))
        elif opcode == OPCODE_PHANTOM_EVASION:
            eva_req = combat_pb2.PhantomEvasionRequest.FromString(payload)
            p = self._get_player_state(eva_req.entity_id or 1001)
            p["anim_state"], p["vx"], p["vy"] = 4, eva_req.evasion_dir_x * 12.0, eva_req.evasion_dir_y * 12.0
            p["x"] += eva_req.evasion_dir_x * 1.5
            p["y"] += eva_req.evasion_dir_y * 1.5
            snap = network_pb2.EntitySnapshot(
                entity_id=eva_req.entity_id or 1001, pos_x=p["x"], pos_y=p["y"], velocity_x=p["vx"],
                velocity_y=p["vy"], current_hp=int(p["hp"]), max_hp=int(p["max_hp"]), animation_state=4,
            )
            await websocket.send(encode_binary_frame(OPCODE_ENTITY_STATE, snap.SerializeToString()))
        elif opcode == OPCODE_COMBAT_DAMAGE_EVENT:
            dmg = combat_pb2.CombatDamageEvent.FromString(payload)
            await websocket.send(encode_binary_frame(OPCODE_COMBAT_DAMAGE_EVENT, dmg.SerializeToString()))
        elif opcode == OPCODE_ENTER_ZONE_REQUEST:
            zone_req = map_zone_pb2.EnterZoneRequest.FromString(payload)
            target_id = "zone_boundless_sanctuary"
            zone_info = map_zone_pb2.ZoneInfoPayload(
                zone_id=target_id, name="Doanh Trại Bến Lưu Đày",
                zone_type=map_zone_pb2.ZoneType.ZONE_TYPE_SANCTUARY, bounds_width=1600.0, bounds_height=1600.0,
            )
            await websocket.send(encode_binary_frame(OPCODE_ZONE_DATA, zone_info.SerializeToString()))
        elif opcode == OPCODE_ZONE_PORTAL_DATA:
            portal = map_zone_pb2.ZonePortalData.FromString(payload)
            await websocket.send(encode_binary_frame(OPCODE_ZONE_PORTAL_DATA, portal.SerializeToString()))
        elif opcode in (OPCODE_CHAT_MESSAGE, 0x0060):
            chat_msg = chat_pb2.ChatMessage.FromString(payload)
            bcast = encode_binary_frame(OPCODE_CHAT_MESSAGE, chat_msg.SerializeToString())
            for client in list(self.clients):
                try:
                    await client.send(bcast)
                except Exception:
                    pass
        elif opcode == OPCODE_AUTH_REQUEST:
            env = network_pb2.PacketEnvelope(
                sequence_number=1, timestamp_ms=int(time.time() * 1000), encrypted_payload=b"AUTH_SIMULATOR_TOKEN_VALID",
            )
            await websocket.send(encode_binary_frame(OPCODE_PACKET_ENVELOPE, env.SerializeToString()))

    async def _handle_json_packet(
        self, websocket: WebSocketServerProtocol, payload: Dict[str, Any]
    ) -> None:
        """Dispatches legacy text JSON frames for backward compatibility."""
        msg_type = payload.get("type", "")
        if msg_type == "ping":
            await websocket.send(json.dumps({
                "type": "pong", "client_ts": payload.get("client_ts", time.time() * 1000), "server_ts": time.time() * 1000,
            }))
        elif msg_type == "skill_intent":
            dmg = max(10, min(1500, int(payload.get("damage", 150))))
            await websocket.send(json.dumps({
                "type": "damage_event", "target_id": payload.get("target_id", "dummy"), "damage": dmg,
                "is_crit": bool(payload.get("is_crit", False)), "verified": True, "timestamp": time.time() * 1000,
            }))
        elif msg_type == "auth_simulator":
            valid = payload.get("token", "").startswith("DEV_SIMULATOR_TOKEN")
            await websocket.send(json.dumps({
                "type": "auth_result", "authorized": valid, "role": "simulator" if valid else "guest",
            }))
        elif msg_type in ("request_zone_data", "enter_zone"):
            zone_id = payload.get("zone_id", "zone_tang_kiem_nhai")
            zone_payload = self.sync_helper.generate_zone_monsters_payload(zone_id)
            await websocket.send(json.dumps(zone_payload))
        elif msg_type == "request_quest_sync":
            player_id = payload.get("player_id", "player_1")
            quest_payload = self.sync_helper.get_player_quests_payload(player_id)
            await websocket.send(json.dumps(quest_payload))
        elif msg_type == "monster_kill_report":
            player_id = payload.get("player_id", "player_1")
            monster_id = payload.get("monster_id", "")
            update_payload = self.sync_helper.process_kill_event(player_id, monster_id)
            await websocket.send(json.dumps(update_payload))
        elif msg_type.startswith("e2e_"):
            import os
            if os.environ.get("FREEEXILE_DEV_MODE") != "1":
                await websocket.send(json.dumps({"type": "error", "msg": "Dev endpoints disabled"}))
                return
                
            if msg_type == "e2e_register_account":
                acc_id = payload.get("account_id")
                self.buyout_engine.register_account(acc_id)
                await websocket.send(json.dumps({"type": "e2e_register_success", "account_id": acc_id}))
            elif msg_type == "e2e_add_item":
                acc_id = payload.get("account_id")
                from server.trade.two_phase_commit import StashItem
                import uuid
                item = StashItem(
                    item_uuid=str(uuid.uuid4()),
                    owner_account_id=acc_id,
                    item_name=payload.get("item_class", "weapon"),
                    asking_price_currency="ThienMenh",
                    asking_price_amount=10
                )
                self.buyout_engine.accounts[acc_id].items[item.item_uuid] = item
                await websocket.send(json.dumps({"type": "e2e_item_added", "item_uuid": item.item_uuid}))
            elif msg_type == "e2e_add_currency":
                acc_id = payload.get("account_id")
                currency = payload.get("currency", "ThienMenh")
                amount = payload.get("amount", 100)
                self.buyout_engine.accounts[acc_id].currencies[currency] += amount
                await websocket.send(json.dumps({"type": "e2e_currency_added"}))
            elif msg_type == "e2e_buyout":
                buyer = payload.get("buyer")
                seller = payload.get("seller")
                item_uuid = payload.get("item_uuid")
                success, msg, tx_id = self.buyout_engine.execute_instant_buyout(
                    buyer, seller, item_uuid, "ThienMenh", payload.get("price", 10)
                )
                await websocket.send(json.dumps({"type": "e2e_buyout_result", "success": success, "msg": msg, "tx_id": tx_id}))
            elif msg_type == "e2e_grant_orb":
                player_id = payload.get("player_id")
                self.orb_service.grant_agent_orb(player_id, "item_agent_orb_tier1", 1)
                await websocket.send(json.dumps({"type": "e2e_orb_granted"}))
            elif msg_type == "e2e_activate_agent":
                player_id = payload.get("player_id")
                from server.agent.llm_provider_client import LLMProviderConfig, LLMProviderType
                from server.agent.agent_decision_core import TacticalStance
                cfg = LLMProviderConfig(provider=LLMProviderType.LOCAL_TACTICAL)
                sess = self.orb_service.activate_delegation(player_id, "item_agent_orb_tier1", TacticalStance.BALANCED, cfg)
                await websocket.send(json.dumps({"type": "e2e_agent_activated", "success": sess is not None}))

    async def _handle_connection(self, websocket: WebSocketServerProtocol) -> None:
        self.clients.add(websocket)
        logger.info("Client connected to WS Gateway Bridge: %s", getattr(websocket, "remote_address", "client"))
        try:
            await websocket.send(json.dumps({
                "type": "welcome", "server_time": time.time(), "authoritative": True, "version": "2.0-Bridge",
            }))
            async for raw_msg in websocket:
                try:
                    if isinstance(raw_msg, (bytes, bytearray)):
                        if len(raw_msg) >= 2:
                            op, body = decode_binary_frame(bytes(raw_msg))
                            await self._handle_binary_packet(websocket, op, body)
                    elif isinstance(raw_msg, str):
                        s = raw_msg.strip()
                        if s.startswith("{"):
                            await self._handle_json_packet(websocket, json.loads(raw_msg))
                        else:
                            raw_b = raw_msg.encode("latin1", errors="ignore")
                            if len(raw_b) >= 2:
                                op, body = decode_binary_frame(raw_b)
                                await self._handle_binary_packet(websocket, op, body)
                except Exception as ex:
                    logger.debug("Error processing frame: %s", ex)
                    continue
        except Exception as e:
            logger.debug("WS connection ended: %s", e)
        finally:
            self.clients.discard(websocket)

    def broadcast_json(self, payload: Dict[str, Any]) -> None:
        """Broadcasts a JSON message to all connected clients."""
        if not self.loop or not self.loop.is_running():
            return
        msg = json.dumps(payload)
        
        async def _bcast():
            for client in list(self.clients):
                try:
                    await client.send(msg)
                except Exception:
                    pass
        
        asyncio.run_coroutine_threadsafe(_bcast(), self.loop)

    async def start_server(self) -> None:
        if websockets is None:
            raise RuntimeError("websockets library not installed")
        self.is_running = True
        self.server = await websockets.serve(self._handle_connection, self.host, self.port)
        logger.info("WsGatewayBridge listening on ws://%s:%d", self.host, self.port)

    async def stop_server(self) -> None:
        self.is_running = False
        for client in list(self.clients):
            try:
                await client.close()
            except Exception:
                pass
        if self.server:
            self.server.close()
            await self.server.wait_closed()
            self.server = None

    def start_background(self) -> None:
        """Starts the WebSocket bridge on a dedicated background event loop thread."""
        if self._thread and self._thread.is_alive():
            return

        def _run_loop():
            self.loop = asyncio.new_event_loop()
            asyncio.set_event_loop(self.loop)
            self.loop.run_until_complete(self.start_server())
            self.loop.run_forever()

        self._thread = threading.Thread(target=_run_loop, daemon=True, name="WsGatewayBridgeThread")
        self._thread.start()
        time.sleep(0.3)

    def stop_background(self) -> None:
        if self.loop and self.loop.is_running():
            self.loop.call_soon_threadsafe(self.loop.stop)
        if self._thread:
            self._thread.join(timeout=1.0)
            self._thread = None


_global_bridge: Optional[WsGatewayBridge] = None


def get_bridge_instance(host: str = "127.0.0.1", port: int = 8080) -> WsGatewayBridge:
    global _global_bridge
    if _global_bridge is None:
        _global_bridge = WsGatewayBridge(host=host, port=port)
    return _global_bridge


if __name__ == "__main__":
    logging.basicConfig(level=logging.INFO)
    bridge = get_bridge_instance()
    asyncio.run(bridge.start_server())
    try:
        asyncio.get_event_loop().run_forever()
    except KeyboardInterrupt:
        print("Stopping WsGatewayBridge...")
