"""
Unit Tests for FreeExile WebSocket Protobuf Gateway Bridge.
Verifies 2-byte opcode binary framing, zero-residual momentum verification,
and full duplex WebSocket communication with Protobuf payloads.
"""

import asyncio
import json
import os
import sys
import unittest

sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../server")))

import websockets
from gateway.ws_gateway_bridge import (
    OPCODE_AUTH_REQUEST,
    OPCODE_CAST_MARTIAL_SKILL,
    OPCODE_CHAT_MESSAGE,
    OPCODE_COMBAT_DAMAGE_EVENT,
    OPCODE_ENTER_ZONE_REQUEST,
    OPCODE_ENTITY_STATE,
    OPCODE_PACKET_ENVELOPE,
    OPCODE_PHANTOM_EVASION,
    OPCODE_PLAYER_MOVE_INPUT,
    OPCODE_ZONE_DATA,
    OPCODE_ZONE_PORTAL_DATA,
    WsGatewayBridge,
    decode_binary_frame,
    encode_binary_frame,
)
from proto import chat_pb2, combat_pb2, map_zone_pb2, network_pb2


class TestWsProtobufGateway(unittest.IsolatedAsyncioTestCase):
    @classmethod
    def setUpClass(cls):
        cls.port = 18080
        cls.host = "127.0.0.1"
        cls.bridge = WsGatewayBridge(host=cls.host, port=cls.port)
        cls.bridge.start_background()

    @classmethod
    def tearDownClass(cls):
        cls.bridge.stop_background()

    def setUp(self):
        self.port = self.__class__.port
        self.host = self.__class__.host
        self.bridge = self.__class__.bridge

    def test_binary_frame_encode_decode(self):
        """Codec correctly serializes and unpacks 2-byte big-endian opcode."""
        raw_payload = b"TEST_BINARY_PROTOBUF_DATA"
        frame = encode_binary_frame(OPCODE_PLAYER_MOVE_INPUT, raw_payload)
        self.assertEqual(len(frame), 2 + len(raw_payload))
        self.assertEqual(frame[0], 0x00)
        self.assertEqual(frame[1], 0x10)

        opcode, decoded_payload = decode_binary_frame(frame)
        self.assertEqual(opcode, OPCODE_PLAYER_MOVE_INPUT)
        self.assertEqual(decoded_payload, raw_payload)

        # Invalid frame too short
        with self.assertRaises(ValueError):
            decode_binary_frame(b"\x00")

    def test_zero_residual_momentum_instant_halt(self):
        """Zero-residual momentum ensures velocity snaps to 0.0 with 0 trailing lerp."""
        entity_id = 5001
        # 1. Moving input
        halted, vx, vy = self.bridge.verify_zero_residual_momentum(entity_id, 1.0, 0.0)
        self.assertFalse(halted)
        self.assertAlmostEqual(vx, 6.0, places=2)
        self.assertAlmostEqual(vy, 0.0, places=2)

        # 2. Key release: dir_x == 0, dir_y == 0
        halted_stop, vx_stop, vy_stop = self.bridge.verify_zero_residual_momentum(entity_id, 0.0, 0.0)
        self.assertTrue(halted_stop)
        self.assertEqual(vx_stop, 0.0)
        self.assertEqual(vy_stop, 0.0)

        # 3. Deadzone threshold: inputs with mag <= 0.05 must also trigger instant halt
        halted_dz, vx_dz, vy_dz = self.bridge.verify_zero_residual_momentum(entity_id, 0.03, 0.04)  # hypot == 0.05
        self.assertTrue(halted_dz)
        self.assertEqual(vx_dz, 0.0)
        self.assertEqual(vy_dz, 0.0)

        # Verify internal state has 0 velocity and idle animation
        state = self.bridge._get_player_state(entity_id)
        self.assertEqual(state["vx"], 0.0)
        self.assertEqual(state["vy"], 0.0)
        self.assertEqual(state["anim_state"], 0)

    async def test_websocket_player_move_input_roundtrip(self):
        """WebSocket roundtrip verifies PlayerMoveInput yields authoritative EntitySnapshot."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            # Welcome message
            welcome = await ws.recv()
            welcome_data = json.loads(welcome)
            self.assertEqual(welcome_data.get("type"), "welcome")

            # 1. Send moving input
            move_msg = network_pb2.PlayerMoveInput(
                entity_id=1001, dir_x=0.0, dir_y=1.0, input_sequence=1
            )
            await ws.send(encode_binary_frame(OPCODE_PLAYER_MOVE_INPUT, move_msg.SerializeToString()))

            resp_frame = await ws.recv()
            self.assertIsInstance(resp_frame, bytes)
            opcode, payload = decode_binary_frame(resp_frame)
            self.assertEqual(opcode, OPCODE_ENTITY_STATE)

            snap = network_pb2.EntitySnapshot.FromString(payload)
            self.assertEqual(snap.entity_id, 1001)
            self.assertAlmostEqual(snap.velocity_y, 6.0, places=1)
            self.assertEqual(snap.animation_state, 1)  # Run
            self.assertEqual(snap.last_processed_input_seq, 1)

            # 2. Send instant stop input (input_x == 0 && input_y == 0)
            stop_msg = network_pb2.PlayerMoveInput(
                entity_id=1001, dir_x=0.0, dir_y=0.0, input_sequence=2
            )
            await ws.send(encode_binary_frame(OPCODE_PLAYER_MOVE_INPUT, stop_msg.SerializeToString()))

            stop_resp = await ws.recv()
            op_stop, payload_stop = decode_binary_frame(stop_resp)
            self.assertEqual(op_stop, OPCODE_ENTITY_STATE)

            snap_stop = network_pb2.EntitySnapshot.FromString(payload_stop)
            self.assertEqual(snap_stop.velocity_x, 0.0)
            self.assertEqual(snap_stop.velocity_y, 0.0)
            self.assertEqual(snap_stop.animation_state, 0)  # Idle
            self.assertEqual(snap_stop.last_processed_input_seq, 2)

    async def test_websocket_combat_skill_and_evasion_roundtrip(self):
        """Verifies CastMartialSkillRequest and PhantomEvasionRequest binary handling."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            await ws.recv()  # welcome

            # Cast skill request
            skill_req = combat_pb2.CastMartialSkillRequest(
                caster_entity_id=1001, skill_id=101, target_x=10.0, target_y=5.0
            )
            await ws.send(encode_binary_frame(OPCODE_CAST_MARTIAL_SKILL, skill_req.SerializeToString()))

            resp = await ws.recv()
            op, payload = decode_binary_frame(resp)
            self.assertEqual(op, OPCODE_COMBAT_DAMAGE_EVENT)
            dmg = combat_pb2.CombatDamageEvent.FromString(payload)
            self.assertEqual(dmg.source_entity_id, 1001)
            self.assertGreater(dmg.raw_damage, 0)

            # Phantom evasion request
            eva_req = combat_pb2.PhantomEvasionRequest(
                entity_id=1001, evasion_dir_x=1.0, evasion_dir_y=0.0
            )
            await ws.send(encode_binary_frame(OPCODE_PHANTOM_EVASION, eva_req.SerializeToString()))

            resp_eva = await ws.recv()
            op_eva, payload_eva = decode_binary_frame(resp_eva)
            self.assertEqual(op_eva, OPCODE_ENTITY_STATE)
            snap_eva = network_pb2.EntitySnapshot.FromString(payload_eva)
            self.assertEqual(snap_eva.animation_state, 4)  # Dodge

    async def test_websocket_zone_and_chat_roundtrip(self):
        """Verifies EnterZoneRequest and ChatMessage binary dispatch."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            await ws.recv()  # welcome

            # Enter zone
            zone_req = map_zone_pb2.EnterZoneRequest(
                player_id="1001", portal_id="portal_sanctuary", player_level=1
            )
            await ws.send(encode_binary_frame(OPCODE_ENTER_ZONE_REQUEST, zone_req.SerializeToString()))

            resp = await ws.recv()
            op, payload = decode_binary_frame(resp)
            self.assertEqual(op, OPCODE_ZONE_DATA)
            zone_info = map_zone_pb2.ZoneInfoPayload.FromString(payload)
            self.assertEqual(zone_info.zone_id, "zone_boundless_sanctuary")
            self.assertEqual(zone_info.bounds_width, 1600.0)

            # Chat message broadcast
            chat_req = chat_pb2.ChatMessage(
                sender_id=1001, sender_name="Exile", raw_content="Savage exile martial path"
            )
            await ws.send(encode_binary_frame(OPCODE_CHAT_MESSAGE, chat_req.SerializeToString()))

            resp_chat = await ws.recv()
            op_chat, payload_chat = decode_binary_frame(resp_chat)
            self.assertEqual(op_chat, OPCODE_CHAT_MESSAGE)
            chat_res = chat_pb2.ChatMessage.FromString(payload_chat)
            self.assertEqual(chat_res.raw_content, "Savage exile martial path")

    async def test_websocket_legacy_json_diagnostics(self):
        """Verifies legacy text JSON frames continue to work without regression."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            await ws.recv()  # welcome

            # JSON ping
            await ws.send(json.dumps({"type": "ping", "client_ts": 9999.0}))
            pong_raw = await ws.recv()
            pong = json.loads(pong_raw)
            self.assertEqual(pong.get("type"), "pong")
            self.assertEqual(pong.get("client_ts"), 9999.0)

            # JSON auth_simulator
            await ws.send(json.dumps({"type": "auth_simulator", "token": "DEV_SIMULATOR_TOKEN_VALID"}))
            auth_raw = await ws.recv()
            auth = json.loads(auth_raw)
            self.assertEqual(auth.get("type"), "auth_result")
            self.assertTrue(auth.get("authorized"))

    async def test_websocket_envelope_portal_auth_roundtrip(self):
        """Verifies PacketEnvelope, ZonePortalData, and AuthRequest binary handling."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            await ws.recv()  # welcome

            # 1. PacketEnvelope (0x0001)
            env_req = network_pb2.PacketEnvelope(sequence_number=42, timestamp_ms=1000, nonce=b"nonce_123456")
            await ws.send(encode_binary_frame(OPCODE_PACKET_ENVELOPE, env_req.SerializeToString()))
            resp_env = await ws.recv()
            op_env, payload_env = decode_binary_frame(resp_env)
            self.assertEqual(op_env, OPCODE_PACKET_ENVELOPE)
            env_ack = network_pb2.PacketEnvelope.FromString(payload_env)
            self.assertEqual(env_ack.sequence_number, 42)
            self.assertEqual(env_ack.nonce, b"nonce_123456")

            # 2. ZonePortalData (0x0021)
            portal = map_zone_pb2.ZonePortalData(portal_id="portal_boss_1", target_zone_id="zone_boss_chamber")
            await ws.send(encode_binary_frame(OPCODE_ZONE_PORTAL_DATA, portal.SerializeToString()))
            resp_p = await ws.recv()
            op_p, payload_p = decode_binary_frame(resp_p)
            self.assertEqual(op_p, OPCODE_ZONE_PORTAL_DATA)
            portal_ack = map_zone_pb2.ZonePortalData.FromString(payload_p)
            self.assertEqual(portal_ack.portal_id, "portal_boss_1")

            # 3. AuthRequest (0x0050)
            await ws.send(encode_binary_frame(OPCODE_AUTH_REQUEST, b""))
            resp_auth = await ws.recv()
            op_auth, payload_auth = decode_binary_frame(resp_auth)
            self.assertEqual(op_auth, OPCODE_PACKET_ENVELOPE)
            auth_ack = network_pb2.PacketEnvelope.FromString(payload_auth)
            self.assertEqual(auth_ack.encrypted_payload, b"AUTH_SIMULATOR_TOKEN_VALID")

    async def test_websocket_zero_residual_momentum_diagonal(self):
        """Verifies diagonal move normalized velocity and instant zero-residual stop."""
        uri = f"ws://{self.host}:{self.port}"
        async with websockets.connect(uri) as ws:
            await ws.recv()  # welcome
            # Diagonal move
            diag_msg = network_pb2.PlayerMoveInput(entity_id=3001, dir_x=0.7071, dir_y=0.7071, input_sequence=1)
            await ws.send(encode_binary_frame(OPCODE_PLAYER_MOVE_INPUT, diag_msg.SerializeToString()))
            resp_diag = await ws.recv()
            _, payload_diag = decode_binary_frame(resp_diag)
            snap_diag = network_pb2.EntitySnapshot.FromString(payload_diag)
            self.assertAlmostEqual(snap_diag.velocity_x, 4.24, places=1)
            self.assertAlmostEqual(snap_diag.velocity_y, 4.24, places=1)
            self.assertEqual(snap_diag.animation_state, 1)

            # Digital brake (0.0, 0.0)
            stop_msg = network_pb2.PlayerMoveInput(entity_id=3001, dir_x=0.0, dir_y=0.0, input_sequence=2)
            await ws.send(encode_binary_frame(OPCODE_PLAYER_MOVE_INPUT, stop_msg.SerializeToString()))
            resp_stop = await ws.recv()
            _, payload_stop = decode_binary_frame(resp_stop)
            snap_stop = network_pb2.EntitySnapshot.FromString(payload_stop)
            self.assertEqual(snap_stop.velocity_x, 0.0)
            self.assertEqual(snap_stop.velocity_y, 0.0)
            self.assertEqual(snap_stop.animation_state, 0)


if __name__ == "__main__":
    unittest.main()
