diff --git a/server/packet_codec.py b/server/packet_codec.py new file mode 100644 index 0000000..1bf0c4a --- /dev/null +++ b/server/packet_codec.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any + +from transport_policy import Delivery, delivery_for_packet_type + +MAX_MESSAGE_BYTES = 64 * 1024 + + +class PacketCodecError(ValueError): + pass + + +@dataclass(frozen=True) +class EncodedPacket: + packet_type: str + payload: bytes + delivery: Delivery + + +def _reject_json_constant(value: str) -> None: + raise PacketCodecError(f"non-finite JSON constant is not allowed: {value}") + + +def encode_packet(packet: dict[str, Any]) -> EncodedPacket: + if not isinstance(packet, dict): + raise PacketCodecError("packet must be a JSON object") + packet_type = packet.get("type") + if not isinstance(packet_type, str) or not packet_type: + raise PacketCodecError("packet type must be a non-empty string") + try: + payload = json.dumps(packet, separators=(",", ":"), allow_nan=False).encode("utf-8") + except (TypeError, ValueError) as error: + raise PacketCodecError(f"packet is not JSON serializable: {error}") from error + if len(payload) > MAX_MESSAGE_BYTES: + raise PacketCodecError("packet exceeds maximum message size") + return EncodedPacket(packet_type, payload, delivery_for_packet_type(packet_type)) + + +def decode_packet(payload: bytes | bytearray | memoryview) -> dict[str, Any]: + raw = bytes(payload) + if len(raw) > MAX_MESSAGE_BYTES: + raise PacketCodecError("packet exceeds maximum message size") + try: + text = raw.decode("utf-8", errors="strict") + except UnicodeDecodeError as error: + raise PacketCodecError("packet is not valid UTF-8") from error + try: + packet = json.loads(text, parse_constant=_reject_json_constant) + except (json.JSONDecodeError, PacketCodecError) as error: + raise PacketCodecError(f"invalid JSON packet: {error}") from error + if not isinstance(packet, dict): + raise PacketCodecError("packet must be a JSON object") + packet_type = packet.get("type") + if not isinstance(packet_type, str) or not packet_type: + raise PacketCodecError("packet type must be a non-empty string") + return packet