"""
Authoritative Gateway Service for FreeExile.
Integrates TCP framing, Apple App Attest authentication, AEAD packet encryption,
RFC 4303 anti-replay protection, touch biometrics anti-bot validation,
and server-authoritative movement with spatial grid synchronization.
"""

import asyncio
import json
import os
import struct
import time
from typing import Dict, Optional, Tuple
from dataclasses import dataclass

from gateway.network_gateway import PacketCodec
from gateway.session_manager import SessionManager, ClientSession
from security.app_attest import AppleAppAttestValidator, AttestationToken
from security.packet_cipher import PacketCipherEngine, AntiReplayWindow
from security.touch_biometrics import TouchBiometricsValidator, TouchPoint
from world.movement_authority import MovementAuthorityEngine, PlayerCharacter
from world.spatial_grid import SpatialGrid, Entity


@dataclass
class ConnectionState:
    session: Optional[ClientSession] = None
    cipher: Optional[PacketCipherEngine] = None
    anti_replay: Optional[AntiReplayWindow] = None
    server_seq: int = 0
    entity_id: Optional[int] = None


class AuthoritativeGatewayService:
    def __init__(self, host: str = "0.0.0.0", port: int = 7777):
        self.host = host
        self.port = port
        self.is_running = False
        self.server: Optional[asyncio.Server] = None
        
        # Subsystems
        self.session_mgr = SessionManager()
        self.attest_validator = AppleAppAttestValidator(expected_bundle_id="com.freeexile.game.ios")
        self.touch_validator = TouchBiometricsValidator()
        self.movement_authority = MovementAuthorityEngine(max_base_speed=6.0)
        self.spatial_grid = SpatialGrid(cell_size=50.0)
        
        self.next_entity_id = 1000

    async def start(self) -> None:
        self.is_running = True
        self.server = await asyncio.start_server(self.handle_client, self.host, self.port)

    async def stop(self) -> None:
        self.is_running = False
        if self.server:
            self.server.close()
            await self.server.wait_closed()

    async def handle_client(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
        conn = ConnectionState()
        buffer = b""

        try:
            while self.is_running:
                chunk = await reader.read(4096)
                if not chunk:
                    break
                buffer += chunk

                while True:
                    payload, buffer = PacketCodec.decode_frame(buffer)
                    if payload is None:
                        break

                    # 1. Unauthenticated state: Expect Handshake
                    if conn.session is None:
                        await self._handle_handshake(payload, conn, writer)
                    else:
                        # 2. Authenticated state: Expect AEAD encrypted packet
                        await self._handle_encrypted_packet(payload, conn, writer)

        except (asyncio.IncompleteReadError, ConnectionResetError):
            pass
        except Exception:
            pass
        finally:
            if conn.session:
                self.session_mgr.remove_session(conn.session.session_id)
                if conn.entity_id:
                    self.spatial_grid.remove_entity(conn.entity_id)
            try:
                writer.close()
                await writer.wait_closed()
            except Exception:
                pass

    async def _handle_handshake(
        self, payload: bytes, conn: ConnectionState, writer: asyncio.StreamWriter
    ) -> None:
        try:
            data = json.loads(payload.decode("utf-8"))
            if data.get("type") != "handshake":
                err_resp = json.dumps({"status": "error", "message": "Expected handshake"}).encode("utf-8")
                writer.write(PacketCodec.encode_frame(err_resp))
                await writer.drain()
                return

            token = AttestationToken(
                device_id=data.get("device_id", ""),
                bundle_id=data.get("bundle_id", ""),
                public_key_hash=data.get("public_key_hash", ""),
                signature=data.get("signature", ""),
                counter=data.get("counter", 1),
                timestamp_ms=data.get("timestamp_ms", int(time.time() * 1000))
            )

            is_valid, reason = self.attest_validator.verify_attestation(token)
            if not is_valid:
                err_resp = json.dumps({"status": "error", "message": reason}).encode("utf-8")
                writer.write(PacketCodec.encode_frame(err_resp))
                await writer.drain()
                return

            # Hardware validated! Setup session and encryption key
            self.next_entity_id += 1
            entity_id = self.next_entity_id
            shared_key = os.urandom(32)

            peer = writer.get_extra_info("peername")
            client_ip = peer[0] if peer else "127.0.0.1"
            session = self.session_mgr.create_session(
                account_id=data.get("account_id", "anon"),
                entity_id=entity_id,
                client_ip=client_ip
            )
            session.is_authenticated = True

            conn.session = session
            conn.entity_id = entity_id
            conn.cipher = PacketCipherEngine(shared_key)
            conn.anti_replay = AntiReplayWindow(window_size=64)
            conn.server_seq = 0

            # Register initial position in movement and grid
            player = PlayerCharacter(entity_id=entity_id, x=0.0, y=0.0, move_speed=6.0)
            self.movement_authority.register_player(player)

            grid_entity = Entity(entity_id=entity_id, x=0.0, y=0.0)
            self.spatial_grid.add_entity(grid_entity)

            ack_resp = json.dumps({
                "status": "ok",
                "session_id": session.session_id,
                "entity_id": entity_id,
                "shared_key": shared_key.hex()
            }).encode("utf-8")

            writer.write(PacketCodec.encode_frame(ack_resp))
            await writer.drain()

        except Exception as ex:
            err_resp = json.dumps({"status": "error", "message": str(ex)}).encode("utf-8")
            writer.write(PacketCodec.encode_frame(err_resp))
            await writer.drain()

    async def _handle_encrypted_packet(
        self, payload: bytes, conn: ConnectionState, writer: asyncio.StreamWriter
    ) -> None:
        # Minimum wire packet: 8 (seq) + 12 (nonce) + 16 (tag) = 36 bytes
        if len(payload) < 36:
            return

        client_seq = struct.unpack(">Q", payload[:8])[0]
        nonce = payload[8:20]
        tag = payload[20:36]
        ciphertext = payload[36:]

        # 1. Anti-Replay Check
        if not conn.anti_replay.check_and_update(client_seq):
            return  # Drop replayed packet silently (RFC 4303 recommendation)

        # 2. Decrypt & Verify AEAD
        try:
            plaintext = conn.cipher.decrypt(ciphertext, tag, nonce, client_seq)
            msg = json.loads(plaintext.decode("utf-8"))
        except Exception:
            return  # Tampered packet, drop

        msg_type = msg.get("type")

        if msg_type == "move":
            # 3. Touch Biometrics Check
            raw_points = msg.get("touch_points", [])
            touch_points = [
                TouchPoint(
                    x=float(p.get("x", 0.0)),
                    y=float(p.get("y", 0.0)),
                    timestamp_ms=float(p.get("timestamp_ms", 0.0)),
                    major_radius=float(p.get("major_radius", 0.0)),
                    force=float(p.get("force", 1.0))
                ) for p in raw_points
            ]

            is_human, trust_score, reason = self.touch_validator.evaluate_touch_sample(touch_points)
            if not is_human:
                # Anomaly detected! Send rejected response
                conn.server_seq += 1
                resp_data = json.dumps({"status": "rejected", "reason": reason}).encode("utf-8")
                resp_cipher, resp_tag, resp_nonce = conn.cipher.encrypt(resp_data, conn.server_seq)
                resp_packet = struct.pack(">Q", conn.server_seq) + resp_nonce + resp_tag + resp_cipher
                writer.write(PacketCodec.encode_frame(resp_packet))
                await writer.drain()
                return

            # 4. Authoritative Movement Validation
            dir_x = float(msg.get("dir_x", 0.0))
            dir_y = float(msg.get("dir_y", 0.0))
            dt = float(msg.get("dt", 0.1))

            success, new_x, new_y = self.movement_authority.process_move_input(
                conn.entity_id, dir_x, dir_y, dt
            )

            if success:
                grid_entity = self.spatial_grid.entities.get(conn.entity_id)
                if grid_entity:
                    self.spatial_grid.update_entity_position(grid_entity, new_x, new_y)

            # 5. Build and send authoritative snapshot
            nearby_entities = self.spatial_grid.get_entities_in_aoi(new_x, new_y, radius_cells=1)
            nearby_ids = [e_id for e_id in nearby_entities if e_id != conn.entity_id]

            snapshot = {
                "type": "snapshot",
                "ack_seq": client_seq,
                "entity_id": conn.entity_id,
                "x": new_x,
                "y": new_y,
                "server_time_ms": int(time.time() * 1000),
                "nearby": list(nearby_ids)
            }

            conn.server_seq += 1
            snap_bytes = json.dumps(snapshot).encode("utf-8")
            snap_cipher, snap_tag, snap_nonce = conn.cipher.encrypt(snap_bytes, conn.server_seq)
            wire_snap = struct.pack(">Q", conn.server_seq) + snap_nonce + snap_tag + snap_cipher
            writer.write(PacketCodec.encode_frame(wire_snap))
            await writer.drain()
