"""Bounded reference codec for the WC3 D3D9 command-stream envelope."""

from __future__ import annotations

from dataclasses import dataclass
from enum import IntEnum, IntFlag
import struct
import time
import zlib


MAGIC = b"W3CS"
VERSION = 1
HEADER = struct.Struct("<4sBBHIIIHHII")
RECORD_HEADER = struct.Struct("<BBHI")
COMPRESSED_HEADER = struct.Struct("<IB3x")
MAX_FRAGMENT = 60 * 1024
MAX_FRAGMENTS = 4096
MAX_FRAME = 4 * 1024 * 1024
MAX_RESOURCE = 64 * 1024 * 1024
MAX_RECORDS = 16_384
FRAME_TTL_SECONDS = 0.100
RESOURCE_TTL_SECONDS = 5.0
MAX_PENDING = 128


class ProtocolError(ValueError):
    """The peer sent a malformed or unsupported command-stream message."""


class Kind(IntEnum):
    HELLO = 1
    RESOURCE = 2
    FRAME = 3
    REPAIR = 4
    ACK = 5
    ERROR = 6


class Flag(IntFlag):
    COMPRESSED = 1
    KEYFRAME = 2
    LAST = 4


class Compression(IntEnum):
    DEFLATE = 1
    ZSTD = 2


@dataclass(frozen=True)
class Envelope:
    kind: Kind
    flags: Flag
    session: int
    sequence: int
    frame: int
    fragment_index: int
    fragment_count: int
    payload: bytes


@dataclass(frozen=True)
class Record:
    opcode: int
    flags: int
    payload: bytes


def _u32(value: int, name: str) -> int:
    if not 0 <= int(value) <= 0xFFFFFFFF:
        raise ProtocolError(f"{name} is outside uint32")
    return int(value)


def encode_envelope(envelope: Envelope) -> bytes:
    payload = bytes(envelope.payload)
    if len(payload) > MAX_FRAGMENT:
        raise ProtocolError("fragment exceeds 60 KiB")
    count = int(envelope.fragment_count)
    index = int(envelope.fragment_index)
    if not 1 <= count <= MAX_FRAGMENTS or not 0 <= index < count:
        raise ProtocolError("invalid fragment coordinates")
    header = HEADER.pack(
        MAGIC,
        VERSION,
        int(Kind(envelope.kind)),
        int(Flag(envelope.flags)),
        _u32(envelope.session, "session"),
        _u32(envelope.sequence, "sequence"),
        _u32(envelope.frame, "frame"),
        index,
        count,
        len(payload),
        zlib.crc32(payload) & 0xFFFFFFFF,
    )
    return header + payload


def decode_envelope(data: bytes) -> Envelope:
    if len(data) < HEADER.size:
        raise ProtocolError("truncated envelope header")
    (magic, version, kind, flags, session, sequence, frame, index, count,
     payload_size, checksum) = HEADER.unpack_from(data)
    if magic != MAGIC or version != VERSION:
        raise ProtocolError("unsupported envelope identity")
    if payload_size > MAX_FRAGMENT or len(data) != HEADER.size + payload_size:
        raise ProtocolError("invalid envelope payload size")
    if not 1 <= count <= MAX_FRAGMENTS or index >= count:
        raise ProtocolError("invalid fragment coordinates")
    payload = data[HEADER.size:]
    if zlib.crc32(payload) & 0xFFFFFFFF != checksum:
        raise ProtocolError("payload checksum mismatch")
    try:
        parsed_kind = Kind(kind)
    except ValueError as error:
        raise ProtocolError("unknown message kind") from error
    return Envelope(parsed_kind, Flag(flags), session, sequence, frame,
                    index, count, payload)


def fragment(kind: Kind, payload: bytes, *, session: int, sequence: int,
             frame: int = 0, flags: Flag = Flag(0)) -> list[bytes]:
    payload = bytes(payload)
    chunks = [payload[offset:offset + MAX_FRAGMENT]
              for offset in range(0, len(payload), MAX_FRAGMENT)] or [b""]
    if len(chunks) > MAX_FRAGMENTS:
        raise ProtocolError("message needs too many fragments")
    result = []
    for index, chunk in enumerate(chunks):
        packet_flags = flags
        if index == len(chunks) - 1:
            packet_flags |= Flag.LAST
        result.append(encode_envelope(Envelope(
            kind, packet_flags, session, sequence + index, frame,
            index, len(chunks), chunk)))
    return result


def encode_records(records: list[Record]) -> bytes:
    if len(records) > MAX_RECORDS:
        raise ProtocolError("too many records")
    output = bytearray()
    for record in records:
        if not 0 <= int(record.opcode) <= 255 or not 0 <= int(record.flags) <= 255:
            raise ProtocolError("record tag is outside uint8")
        payload = bytes(record.payload)
        output.extend(RECORD_HEADER.pack(record.opcode, record.flags, 0,
                                         len(payload)))
        output.extend(payload)
    return bytes(output)


def decode_records(payload: bytes) -> list[Record]:
    records = []
    offset = 0
    while offset < len(payload):
        if len(records) >= MAX_RECORDS or len(payload) - offset < RECORD_HEADER.size:
            raise ProtocolError("invalid record stream")
        opcode, flags, reserved, size = RECORD_HEADER.unpack_from(payload, offset)
        offset += RECORD_HEADER.size
        if reserved or size > len(payload) - offset:
            raise ProtocolError("invalid record header")
        records.append(Record(opcode, flags, payload[offset:offset + size]))
        offset += size
    return records


def compress_payload(payload: bytes, codec: Compression = Compression.DEFLATE,
                     level: int = 1) -> tuple[Flag, bytes]:
    """Compress one bounded logical message when doing so saves wire bytes."""
    payload = bytes(payload)
    if codec is not Compression.DEFLATE:
        raise ProtocolError("compression codec is not installed")
    compressed = zlib.compress(payload, level=level)
    framed = COMPRESSED_HEADER.pack(len(payload), int(codec)) + compressed
    if len(framed) >= len(payload):
        return Flag(0), payload
    return Flag.COMPRESSED, framed


def decompress_payload(payload: bytes, *, kind: Kind,
                       flags: Flag) -> bytes:
    if not flags & Flag.COMPRESSED:
        return bytes(payload)
    if len(payload) < COMPRESSED_HEADER.size:
        raise ProtocolError("truncated compressed payload")
    expected, codec_value = COMPRESSED_HEADER.unpack_from(payload)
    limit = MAX_FRAME if kind is Kind.FRAME else MAX_RESOURCE
    if expected > limit:
        raise ProtocolError("uncompressed payload exceeds limit")
    try:
        codec = Compression(codec_value)
    except ValueError as error:
        raise ProtocolError("unknown compression codec") from error
    if codec is not Compression.DEFLATE:
        raise ProtocolError("compression codec is not installed")
    try:
        result = zlib.decompress(payload[COMPRESSED_HEADER.size:])
    except zlib.error as error:
        raise ProtocolError("invalid compressed payload") from error
    if len(result) != expected:
        raise ProtocolError("uncompressed payload size mismatch")
    return result


class Reassembler:
    """Reassemble out-of-order fragments with strict memory and time bounds."""

    def __init__(self, clock=time.monotonic):
        self._clock = clock
        self._pending = {}

    def _limit(self, kind: Kind) -> int:
        return MAX_FRAME if kind is Kind.FRAME else MAX_RESOURCE

    def expire(self) -> None:
        now = self._clock()
        self._pending = {
            key: value for key, value in self._pending.items()
            if now - value[0] <= (
                FRAME_TTL_SECONDS if key[1] is Kind.FRAME
                else RESOURCE_TTL_SECONDS)
        }

    def push(self, packet: bytes) -> tuple[Envelope, bytes] | None:
        envelope = decode_envelope(packet)
        self.expire()
        message_sequence = envelope.sequence - envelope.fragment_index
        key = (envelope.session, envelope.kind, envelope.frame, message_sequence)
        pending = self._pending.get(key)
        if pending is None:
            if len(self._pending) >= MAX_PENDING:
                raise ProtocolError("too many incomplete messages")
            pending = (self._clock(), envelope.fragment_count, {}, envelope.flags)
            self._pending[key] = pending
        created, count, pieces, combined_flags = pending
        if count != envelope.fragment_count:
            del self._pending[key]
            raise ProtocolError("fragment count changed within message")
        previous = pieces.get(envelope.fragment_index)
        if previous is not None and previous != envelope.payload:
            del self._pending[key]
            raise ProtocolError("conflicting duplicate fragment")
        pieces[envelope.fragment_index] = envelope.payload
        total = sum(len(piece) for piece in pieces.values())
        if total > self._limit(envelope.kind):
            del self._pending[key]
            raise ProtocolError("reassembled message exceeds limit")
        self._pending[key] = (created, count, pieces,
                              Flag(combined_flags | envelope.flags))
        if len(pieces) != count:
            return None
        payload = b"".join(pieces[index] for index in range(count))
        del self._pending[key]
        complete = Envelope(
            envelope.kind, Flag(combined_flags | envelope.flags),
            envelope.session, envelope.sequence, envelope.frame, 0, 1, payload)
        return complete, payload
