"""
Dynamic Packet Cipher, Timestamp Gating & Anti-Replay Guard for FreeExile.
Implements Authenticated Encryption with Associated Data (AEAD),
RFC-compliant Sliding Window Anti-Replay, and Sub-Millisecond Nonce Caching (< 1ms).
"""

from __future__ import annotations
import hmac
import hashlib
import os
import struct
import time
import secrets
from typing import Tuple, Dict, Optional
from dataclasses import dataclass


@dataclass(slots=True, frozen=True)
class SecureNetworkPacket:
    packet_id: str
    seq_num: int
    timestamp_ms: int
    nonce: str
    payload: bytes
    signature: str


class PacketCipherEngine:
    """AEAD-style packet encryption with sequence-derived keystream and HMAC-SHA256."""

    def __init__(self, shared_key: bytes):
        if len(shared_key) < 32:
            shared_key = hashlib.sha256(shared_key).digest()
        self.enc_key = shared_key[:16]
        self.mac_key = shared_key[16:32]

    def encrypt(self, plaintext: bytes, seq_num: int) -> Tuple[bytes, bytes, bytes]:
        """Encrypts payload with sequence-derived keystream and calculates HMAC authentication tag."""
        if seq_num < 0:
            raise ValueError("Sequence number must be non-negative.")
        nonce = os.urandom(12)
        keystream_seed = self.enc_key + nonce + struct.pack(">Q", seq_num)
        
        keystream = bytearray()
        block_idx = 0
        while len(keystream) < len(plaintext):
            keystream.extend(hashlib.sha256(keystream_seed + struct.pack(">I", block_idx)).digest())
            block_idx += 1
        
        ciphertext = bytes(p ^ k for p, k in zip(plaintext, keystream[:len(plaintext)]))
        tag_data = struct.pack(">Q", seq_num) + nonce + ciphertext
        tag = hmac.new(self.mac_key, tag_data, hashlib.sha256).digest()[:16]

        return ciphertext, tag, nonce

    def decrypt(self, ciphertext: bytes, tag: bytes, nonce: bytes, seq_num: int) -> bytes:
        """Verifies authenticity tag before decrypting. Raises ValueError if tampering detected."""
        if seq_num < 0:
            raise ValueError("Sequence number must be non-negative.")
        if len(tag) != 16:
            raise ValueError("Packet authentication tag verification failed. Invalid tag length.")
        tag_data = struct.pack(">Q", seq_num) + nonce + ciphertext
        expected_tag = hmac.new(self.mac_key, tag_data, hashlib.sha256).digest()[:16]

        if not hmac.compare_digest(tag, expected_tag):
            raise ValueError("Packet authentication tag verification failed. Tampering detected.")

        keystream_seed = self.enc_key + nonce + struct.pack(">Q", seq_num)
        keystream = bytearray()
        block_idx = 0
        while len(keystream) < len(ciphertext):
            keystream.extend(hashlib.sha256(keystream_seed + struct.pack(">I", block_idx)).digest())
            block_idx += 1

        return bytes(c ^ k for c, k in zip(ciphertext, keystream[:len(ciphertext)]))


class AntiReplayWindow:
    """Sliding window anti-replay filter (RFC 4303 IPsec Anti-Replay model)."""

    def __init__(self, window_size: int = 64):
        self.window_size = window_size
        self.highest_seq: int = 0
        self.bitmap: int = 0

    def check_and_update(self, seq_num: int) -> bool:
        """Validates seq_num against sliding window. Returns True if accepted, False if replayed."""
        if seq_num <= 0:
            return False

        if self.highest_seq == 0:
            self.highest_seq = seq_num
            self.bitmap = 1
            return True

        if seq_num > self.highest_seq:
            diff = seq_num - self.highest_seq
            if diff < self.window_size:
                self.bitmap = (self.bitmap << diff) | 1
            else:
                self.bitmap = 1
            self.highest_seq = seq_num
            return True

        diff = self.highest_seq - seq_num
        if diff >= self.window_size:
            return False

        mask = 1 << diff
        if self.bitmap & mask:
            return False

        self.bitmap |= mask
        return True


class TimestampGatedPacketGuard:
    """
    High-Performance Timestamp Gating & Nonce Cache Guard.
    Enforces maximum timestamp drift (< 3000ms), single-use nonces, and HMAC integrity.
    Rejects replay attacks in < 1ms.
    """

    def __init__(
        self,
        shared_key: bytes,
        max_drift_ms: int = 3000,
        window_size: int = 64,
    ) -> None:
        self.shared_key = shared_key
        self.max_drift_ms = max_drift_ms
        self.replay_window = AntiReplayWindow(window_size=window_size)
        # nonce -> timestamp_ms
        self.nonce_cache: Dict[str, int] = {}

    def create_packet(
        self,
        payload: bytes,
        seq_num: int,
        timestamp_ms: Optional[int] = None,
        nonce: Optional[str] = None,
    ) -> SecureNetworkPacket:
        """Creates a signed, timestamped, nonce-protected network packet."""
        ts = int(time.time() * 1000) if timestamp_ms is None else timestamp_ms
        n = secrets.token_hex(16) if nonce is None else nonce
        packet_id = f"pkt_{seq_num}_{n[:8]}"
        sign_payload = f"{seq_num}:{ts}:{n}".encode("utf-8") + payload
        sig = hmac.new(self.shared_key, sign_payload, hashlib.sha256).hexdigest()

        return SecureNetworkPacket(
            packet_id=packet_id,
            seq_num=seq_num,
            timestamp_ms=ts,
            nonce=n,
            payload=payload,
            signature=sig,
        )

    def verify_and_process_packet(
        self,
        packet: SecureNetworkPacket,
        server_now_ms: Optional[int] = None,
    ) -> Tuple[bool, str, float]:
        """
        Validates timestamp drift, HMAC signature, nonce uniqueness, and sequence number.
        Returns: (is_valid: bool, status_reason: str, processing_latency_ms: float)
        Guaranteed to reject replayed or invalid packets in < 1.0ms.
        """
        t0 = time.perf_counter()
        now_ms = int(time.time() * 1000) if server_now_ms is None else server_now_ms

        # 1. Timestamp drift gating
        drift = abs(now_ms - packet.timestamp_ms)
        if drift > self.max_drift_ms:
            elapsed_ms = (time.perf_counter() - t0) * 1000.0
            return False, f"TIMESTAMP_DRIFT_EXCEEDED ({drift}ms > {self.max_drift_ms}ms)", elapsed_ms

        # 2. Cryptographic HMAC signature check
        sign_payload = f"{packet.seq_num}:{packet.timestamp_ms}:{packet.nonce}".encode("utf-8") + packet.payload
        expected_sig = hmac.new(self.shared_key, sign_payload, hashlib.sha256).hexdigest()
        if not hmac.compare_digest(packet.signature, expected_sig):
            elapsed_ms = (time.perf_counter() - t0) * 1000.0
            return False, "INVALID_HMAC_SIGNATURE", elapsed_ms

        # 3. Nonce uniqueness cache (Single-use token defense)
        if packet.nonce in self.nonce_cache:
            elapsed_ms = (time.perf_counter() - t0) * 1000.0
            return False, "REPLAY_DUPLICATE_NONCE", elapsed_ms

        # 4. Anti-replay sliding sequence window
        if not self.replay_window.check_and_update(packet.seq_num):
            elapsed_ms = (time.perf_counter() - t0) * 1000.0
            return False, "REPLAY_SEQUENCE_NUMBER", elapsed_ms

        # 5. Cache nonce and check pruning
        self.nonce_cache[packet.nonce] = packet.timestamp_ms
        if len(self.nonce_cache) > 50000:
            self.prune_expired_nonces(now_ms)

        elapsed_ms = (time.perf_counter() - t0) * 1000.0
        return True, "PACKET_ACCEPTED", elapsed_ms

    def prune_expired_nonces(self, server_now_ms: Optional[int] = None) -> int:
        """Prunes cached nonces older than 2x max_drift_ms to keep memory lean."""
        now_ms = int(time.time() * 1000) if server_now_ms is None else server_now_ms
        cutoff = now_ms - (self.max_drift_ms * 2)
        expired = [n for n, ts in self.nonce_cache.items() if ts < cutoff]
        for n in expired:
            del self.nonce_cache[n]
        return len(expired)
