"""
Tier 1: Feature Coverage — Protobuf Schemas, Opcode Framing & 30Hz Reconciliation.
Covers Features 8, 9, 10, 11 from PROJECT.md Feature Inventory:
- Feature 8: Protobufjs ES Module Compilation & Schemas
- Feature 9: 2-Byte Opcode Network Framing
- Feature 10: WebSocket Bridge Port 8080
- Feature 11: 30Hz Server Reconciliation
"""

from __future__ import annotations

import json
import math
import struct
import time
from typing import Any, Dict, List, Tuple
import pytest

from server.proto import (
    auth_pb2,
    combat_pb2,
    map_zone_pb2,
    network_pb2,
)
from tests.e2e_cocos.conftest import (
    OPCODES,
    decode_opcode_frame,
    encode_opcode_frame,
)


class TestTier1ProtobufSchemas:
    """Feature 8: Protobuf Schemas & Wire Codec."""

    def test_t1_f8_player_move_input_serialization(self) -> None:
        msg = network_pb2.PlayerMoveInput(
            entity_id=101,
            dir_x=0.7071,
            dir_y=0.7071,
            input_sequence=42,
        )
        msg.biometrics.major_radius = 20.0
        msg.biometrics.force = 1.0
        msg.biometrics.micro_tremor_hz = 10.5
        msg.biometrics.trajectory_curvature = 0.8
        msg.biometrics.device_type = 1

        binary_data = msg.SerializeToString()
        assert len(binary_data) > 0

        # Decode and verify exact field parity
        decoded = network_pb2.PlayerMoveInput()
        decoded.ParseFromString(binary_data)
        assert decoded.entity_id == 101
        assert decoded.input_sequence == 42
        assert math.isclose(decoded.dir_x, 0.7071, abs_tol=1e-4)
        assert decoded.biometrics.device_type == 1

    def test_t1_f8_entity_snapshot_reconciliation_fields(self) -> None:
        snap = network_pb2.EntitySnapshot(
            entity_id=101,
            pos_x=12.5,
            pos_y=-34.2,
            velocity_x=5.5,
            velocity_y=0.0,
            current_hp=100,
            max_hp=100,
            animation_state=1,  # RUN
            last_processed_input_seq=42,
        )
        bin_snap = snap.SerializeToString()
        decoded = network_pb2.EntitySnapshot()
        decoded.ParseFromString(bin_snap)
        assert decoded.last_processed_input_seq == 42
        assert math.isclose(decoded.pos_x, 12.5, abs_tol=1e-4)

    def test_t1_f8_world_state_sync_30hz_broadcast(self) -> None:
        world_sync = network_pb2.WorldStateSync(server_tick=1200)
        e1 = world_sync.entities.add(entity_id=1, pos_x=0.0, pos_y=0.0)
        e2 = world_sync.entities.add(entity_id=2, pos_x=5.0, pos_y=5.0)
        bin_sync = world_sync.SerializeToString()
        decoded = network_pb2.WorldStateSync()
        decoded.ParseFromString(bin_sync)
        assert decoded.server_tick == 1200
        assert len(decoded.entities) == 2

    def test_t1_f8_cast_martial_skill_request(self) -> None:
        skill_req = combat_pb2.CastMartialSkillRequest(
            caster_entity_id=101,
            skill_id=1001,
            target_x=15.0,
            target_y=12.0,
            active_weapon_set=1,
            client_timestamp_ms=1727941234567,
            auto_weapon_swapped=False,
        )
        bin_data = skill_req.SerializeToString()
        decoded = combat_pb2.CastMartialSkillRequest()
        decoded.ParseFromString(bin_data)
        assert decoded.skill_id == 1001
        assert decoded.target_x == 15.0

    def test_t1_f8_combat_damage_event_broadcast(self) -> None:
        dmg_event = combat_pb2.CombatDamageEvent(
            source_entity_id=101,
            target_entity_id=202,
            raw_damage=350,
            mitigated_damage=280,
            is_critical=True,
            target_evaded=False,
            element=combat_pb2.FiveElementsType.ELEMENT_FIRE_HOA,
        )
        bin_data = dmg_event.SerializeToString()
        decoded = combat_pb2.CombatDamageEvent()
        decoded.ParseFromString(bin_data)
        assert decoded.mitigated_damage == 280
        assert decoded.is_critical is True
        assert decoded.element == combat_pb2.FiveElementsType.ELEMENT_FIRE_HOA

    def test_t1_f8_zone_info_payload_schema(self) -> None:
        zone = map_zone_pb2.ZoneInfoPayload(
            zone_id="zone_tang_kiem_nhai",
            name="Bờ Đá Tàn Xương",
            zone_type=map_zone_pb2.ZoneType.ZONE_TYPE_OPEN_WORLD,
            min_level=1,
            bounds_width=1600.0,
            bounds_height=1600.0,
        )
        wp = zone.waypoints.add(waypoint_id="wp_start", x=0.0, y=0.0, is_unlocked=True)
        bin_data = zone.SerializeToString()
        decoded = map_zone_pb2.ZoneInfoPayload()
        decoded.ParseFromString(bin_data)
        assert decoded.bounds_width == 1600.0
        assert len(decoded.waypoints) == 1


class TestTier1OpcodeFraming:
    """Feature 9: 2-Byte Opcode Network Framing."""

    def test_t1_f9_encode_and_decode_player_move_opcode(self) -> None:
        move_msg = network_pb2.PlayerMoveInput(entity_id=101, dir_x=1.0, dir_y=0.0)
        bin_payload = move_msg.SerializeToString()
        frame = encode_opcode_frame(OPCODES["PLAYER_MOVE"], bin_payload)
        opcode, payload = decode_opcode_frame(frame)
        assert opcode == 0x0010
        decoded_move = network_pb2.PlayerMoveInput()
        decoded_move.ParseFromString(payload)
        assert decoded_move.entity_id == 101

    def test_t1_f9_encode_world_sync_opcode(self) -> None:
        sync_msg = network_pb2.WorldStateSync(server_tick=999)
        frame = encode_opcode_frame(OPCODES["WORLD_SYNC"], sync_msg.SerializeToString())
        opcode, payload = decode_opcode_frame(frame)
        assert opcode == 0x0011
        decoded_sync = network_pb2.WorldStateSync()
        decoded_sync.ParseFromString(payload)
        assert decoded_sync.server_tick == 999

    def test_t1_f9_big_endian_byte_ordering(self) -> None:
        # 0x0010 in big-endian bytes is b'\x00\x10'
        frame = encode_opcode_frame(0x0010, b"")
        assert frame == b"\x00\x10"

    def test_t1_f9_combat_skill_opcode(self) -> None:
        frame = encode_opcode_frame(OPCODES["CAST_SKILL"], b"payload")
        opcode, payload = decode_opcode_frame(frame)
        assert opcode == 0x0020
        assert payload == b"payload"

    def test_t1_f9_combat_damage_opcode(self) -> None:
        frame = encode_opcode_frame(OPCODES["COMBAT_DAMAGE"], b"damage_data")
        opcode, payload = decode_opcode_frame(frame)
        assert opcode == 0x0022
        assert payload == b"damage_data"

    def test_t1_f9_zone_opcode(self) -> None:
        frame = encode_opcode_frame(OPCODES["ENTER_ZONE_REQ"], b"zone_request")
        opcode, payload = decode_opcode_frame(frame)
        assert opcode == 0x0030
        assert payload == b"zone_request"


class TestTier1WebSocketBridgePort8080:
    """Feature 10: WebSocket Bridge Port 8080."""

    def test_t1_f10_handshake_welcome_packet_structure(self) -> None:
        welcome_json = {
            "type": "welcome",
            "server_time": time.time(),
            "authoritative": True,
            "version": "2.0-Bridge",
        }
        raw_msg = json.dumps(welcome_json)
        parsed = json.loads(raw_msg)
        assert parsed["type"] == "welcome"
        assert parsed["authoritative"] is True
        assert parsed["version"] == "2.0-Bridge"

    def test_t1_f10_ping_pong_rtt_calculation(self) -> None:
        send_time_ms = 1000.0
        receive_time_ms = 1014.5
        rtt_ms = receive_time_ms - send_time_ms
        assert rtt_ms == 14.5
        assert rtt_ms <= 15.0  # Within target network benchmark

    def test_t1_f10_simulator_token_authentication(self) -> None:
        valid_token = "DEV_SIMULATOR_TOKEN_LOCAL"
        invalid_token = "RANDOM_UNAUTHORIZED_TOKEN"
        assert valid_token.startswith("DEV_SIMULATOR_TOKEN")
        assert not invalid_token.startswith("DEV_SIMULATOR_TOKEN")

    def test_t1_f10_dual_mode_routing_detection(self) -> None:
        json_frame = b'{"type":"ping","client_ts":1234}'
        binary_frame = b"\x00\x10\x08\x65"
        # First byte check: '{' = 0x7b indicates JSON; non-ASCII indicates binary
        assert json_frame[0] == ord("{")
        assert binary_frame[0] != ord("{")

    def test_t1_f10_reconnection_backoff_algorithm(self) -> None:
        max_delay = 10.0
        # Exponential backoff: min(max_delay, 1.0 * (1.5 ^ attempts))
        delays = [min(max_delay, 1.0 * (1.5 ** i)) for i in range(5)]
        assert delays[0] == 1.0
        assert delays[1] == 1.5
        assert delays[2] == 2.25
        assert delays[4] < max_delay

    def test_t1_f10_clean_connection_lifecycle_states(self) -> None:
        states = ["DISCONNECTED", "CONNECTING", "CONNECTED", "AUTHENTICATED"]
        curr_state = states[0]
        # Transition on open
        curr_state = states[2]
        # Transition on auth
        curr_state = states[3]
        assert curr_state == "AUTHENTICATED"


class TestTier1ServerReconciliation:
    """Feature 11: 30Hz Server Reconciliation."""

    def test_t1_f11_input_sequence_tracking_and_discard(self) -> None:
        pending_inputs = [
            {"seq": 1, "dx": 1.0, "dy": 0.0},
            {"seq": 2, "dx": 1.0, "dy": 0.0},
            {"seq": 3, "dx": 1.0, "dy": 0.0},
            {"seq": 4, "dx": 1.0, "dy": 0.0},
        ]
        ack_seq = 2
        # Discard inputs with seq <= ack_seq
        remaining = [inp for inp in pending_inputs if inp["seq"] > ack_seq]
        assert len(remaining) == 2
        assert remaining[0]["seq"] == 3

    def test_t1_f11_noise_tolerance_deadzone_below_5cm(self) -> None:
        pred_x, pred_y = 10.0, 5.0
        auth_x, auth_y = 10.02, 5.03
        error_dist = math.hypot(auth_x - pred_x, auth_y - pred_y)
        # Error is ~0.036m <= 0.05m -> treated as float noise, no correction applied
        assert error_dist < 0.05

    def test_t1_f11_exponential_damping_correction_25_percent(self) -> None:
        pred_x = 10.0
        auth_x = 10.40  # error = +0.40m
        error = auth_x - pred_x
        # 0.25 correction applied
        corrected_x = pred_x + 0.25 * error
        assert math.isclose(corrected_x, 10.10, abs_tol=1e-4)

    def test_t1_f11_hard_rubberband_threshold_above_2m(self) -> None:
        pred_x = 10.0
        auth_x = 13.0  # error = 3.0m > 2.0m -> wall collision rejection
        error = abs(auth_x - pred_x)
        # Must snap immediately to authoritative truth
        snapped_x = auth_x if error > 2.0 else pred_x
        assert snapped_x == auth_x

    def test_t1_f11_remote_entity_snapshot_interpolation_lerp(self) -> None:
        t1, pos1 = 0.0, 10.0
        t2, pos2 = 0.0666, 12.0
        t_render = 0.0333
        alpha = (t_render - t1) / (t2 - t1)
        interp_pos = pos1 + alpha * (pos2 - pos1)
        assert math.isclose(interp_pos, 11.0, abs_tol=1e-2)

    def test_t1_f11_remote_ring_buffer_30_snapshots_cap(self) -> None:
        buffer_cap = 30
        snap_buffer = []
        for i in range(50):
            snap_buffer.append(i)
            if len(snap_buffer) > buffer_cap:
                snap_buffer.pop(0)
        assert len(snap_buffer) == 30
        assert snap_buffer[-1] == 49
