"""CMVR edge QUIC gateway. Control stream records are protobuf messages prefixed by a four-byte, network-byte-order unsigned length. """ import asyncio import json import os import signal import struct import time import uuid from base64 import b64encode from pathlib import Path from aioquic.asyncio import QuicConnectionProtocol from aioquic.asyncio.server import QuicServer from aioquic.quic.configuration import QuicConfiguration from aioquic.quic.events import ( ConnectionTerminated, DatagramFrameReceived, HandshakeCompleted, StreamDataReceived, ) from google.protobuf.json_format import MessageToDict from generated import quic_edge_pb2 PROTOCOL_VERSION = 1 def log(event: str, **fields): print(json.dumps({"source": "quic", "event": event, **fields}, ensure_ascii=False), flush=True) MAX_CONTROL_FRAME_BYTES = int(os.getenv("QUIC_MAX_CONTROL_FRAME_BYTES", "1048576")) DATAGRAM_HEADER_BYTES = 64 DATAGRAM_MAGIC = b"CMQD" MEDIA_REASSEMBLY_TIMEOUT_SECONDS = 5 def pop_delimited(buffer: bytearray): if len(buffer) < 4: return None length = int.from_bytes(buffer[:4], "big") if length == 0 or length > MAX_CONTROL_FRAME_BYTES: raise ValueError(f"invalid control frame length: {length}") if len(buffer) < 4 + length: return None payload = bytes(buffer[4:4 + length]) del buffer[:4 + length] return payload class CmvrQuicProtocol(QuicConnectionProtocol): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.buffers = {} self.sessions = {} self.out_sequence = 0 self.media_frames = {} def peer(self): paths = getattr(self._quic, "_network_paths", []) address = paths[0].addr if paths else ("unknown", 0) return {"ip": str(address[0]), "port": int(address[1])} def quic_event_received(self, event): if isinstance(event, HandshakeCompleted): log("connected", peer=self.peer(), alpn=event.alpn_protocol) elif isinstance(event, StreamDataReceived): self.receive_stream(event.stream_id, event.data, event.end_stream) elif isinstance(event, DatagramFrameReceived): self.receive_datagram(event.data) elif isinstance(event, ConnectionTerminated): log("disconnected", peer=self.peer(), error_code=event.error_code, reason=event.reason_phrase) def receive_stream(self, stream_id: int, data: bytes, end_stream: bool): buffer = self.buffers.setdefault(stream_id, bytearray()) buffer.extend(data) try: while payload := pop_delimited(buffer): envelope = quic_edge_pb2.EdgeControlEnvelope() envelope.ParseFromString(payload) self.handle_envelope(stream_id, envelope) except Exception as error: log("protocol_error", peer=self.peer(), stream_id=stream_id, message=str(error)) self.send_error(stream_id, 1, str(error), fatal=True) if end_stream: if buffer: log("truncated_message", peer=self.peer(), stream_id=stream_id, remaining_bytes=len(buffer)) self.buffers.pop(stream_id, None) def handle_envelope(self, stream_id, envelope): payload_type = envelope.WhichOneof("payload") message = getattr(envelope, payload_type) if payload_type else None log( "control_message", peer=self.peer(), stream_id=stream_id, protocol_version=envelope.protocol_version, message_sequence=str(envelope.message_sequence), payload_type=payload_type, payload=MessageToDict(message, preserving_proto_field_name=True) if message else None, ) if envelope.protocol_version != PROTOCOL_VERSION: self.send_error(stream_id, 2, f"unsupported protocol version: {envelope.protocol_version}", envelope.message_sequence, True) return if payload_type == "node_register_request": request = envelope.node_register_request session_id = str(uuid.uuid4()) self.sessions[request.node.node_id] = session_id response = quic_edge_pb2.EdgeControlEnvelope(protocol_version=PROTOCOL_VERSION) response.node_register_response.CopyFrom(quic_edge_pb2.NodeRegisterResponse( accepted=True, session_id=session_id, message="registered", heartbeat_interval_ms=int(os.getenv("QUIC_HEARTBEAT_INTERVAL_MS", "5000")), observed_source_ip=self.peer()["ip"], )) self.send_envelope(stream_id, response) elif payload_type == "node_heartbeat": heartbeat = envelope.node_heartbeat expected = self.sessions.get(heartbeat.node_id) accepted = bool(expected and expected == heartbeat.session_id) response = quic_edge_pb2.EdgeControlEnvelope(protocol_version=PROTOCOL_VERSION) response.node_heartbeat_ack.CopyFrom(quic_edge_pb2.NodeHeartbeatAck( accepted=accepted, acknowledged_sequence=heartbeat.sequence, message="ok" if accepted else "unknown or mismatched session", server_time_unix_ms=int(time.time() * 1000), observed_source_ip=self.peer()["ip"], session_id=expected or "", )) self.send_envelope(stream_id, response) def receive_datagram(self, data: bytes): now = time.monotonic() for key, frame in list(self.media_frames.items()): if now - frame["updated_at"] > MEDIA_REASSEMBLY_TIMEOUT_SECONDS: del self.media_frames[key] if len(data) < DATAGRAM_HEADER_BYTES or data[:4] != DATAGRAM_MAGIC: log("invalid_datagram", peer=self.peer(), size=len(data), message="short packet or magic mismatch") return values = struct.unpack(">4sBBHHHHHIIQQQQII", data[:DATAGRAM_HEADER_BYTES]) (_, version, kind, flags, header_size, fragment_index, fragment_count, payload_size, track_id, generation_token, session_epoch, packet_sequence, frame_sequence, capture_timestamp_us, frame_size, fragment_offset) = values payload = data[DATAGRAM_HEADER_BYTES:] if (version != PROTOCOL_VERSION or header_size != DATAGRAM_HEADER_BYTES or kind not in (1, 2) or fragment_count == 0 or fragment_index >= fragment_count or payload_size != len(payload) or frame_size == 0 or fragment_offset + payload_size > frame_size): log("invalid_datagram", peer=self.peer(), size=len(data), message="invalid CMQD header fields") return header = { "kind": "video" if kind == 1 else "audio", "flags": flags, "key_frame": bool(flags & 1), "discontinuity": bool(flags & 4), "fragment_index": fragment_index, "fragment_count": fragment_count, "payload_size": payload_size, "track_id": track_id, "codec_generation_token": generation_token, "session_epoch": str(session_epoch), "packet_sequence": str(packet_sequence), "frame_sequence": str(frame_sequence), "capture_timestamp_us": str(capture_timestamp_us), "frame_size": frame_size, "fragment_offset": fragment_offset, } log("media_fragment", peer=self.peer(), **header) key = (session_epoch, track_id, frame_sequence) frame = self.media_frames.setdefault(key, { "data": bytearray(frame_size), "received": set(), "header": header, "fragment_count": fragment_count, "updated_at": now, }) if frame["fragment_count"] != fragment_count or len(frame["data"]) != frame_size: del self.media_frames[key] log("invalid_datagram", peer=self.peer(), message="inconsistent fragments for frame", **header) return frame["data"][fragment_offset:fragment_offset + payload_size] = payload frame["received"].add(fragment_index) frame["updated_at"] = now if len(frame["received"]) == fragment_count: log("media_frame", peer=self.peer(), **frame["header"], data_base64=b64encode(frame["data"]).decode()) del self.media_frames[key] def send_envelope(self, stream_id, envelope): self.out_sequence += 1 envelope.message_sequence = self.out_sequence data = envelope.SerializeToString() self._quic.send_stream_data(stream_id, len(data).to_bytes(4, "big") + data, end_stream=False) self.transmit() def send_error(self, stream_id, code, message, related_sequence=0, fatal=False): envelope = quic_edge_pb2.EdgeControlEnvelope(protocol_version=PROTOCOL_VERSION) envelope.protocol_error.CopyFrom(quic_edge_pb2.ProtocolError( code=code, message=message, related_message_sequence=related_sequence, fatal=fatal )) self.send_envelope(stream_id, envelope) class LoggingQuicServer(QuicServer): """Expose packet arrival even when QUIC/TLS negotiation cannot complete.""" def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self._last_packet_log = {} def datagram_received(self, data, addr): now = time.monotonic() if now - self._last_packet_log.get(addr, 0) >= 5: log("udp_packet_received", peer={"ip": str(addr[0]), "port": int(addr[1])}, size=len(data)) self._last_packet_log[addr] = now super().datagram_received(data, addr) async def main(): host = os.getenv("QUIC_HOST", "0.0.0.0") port = int(os.getenv("QUIC_PORT", "5175")) cert = Path(os.getenv("QUIC_CERT_FILE", "server/certs/quic-gateway.crt")) key = Path(os.getenv("QUIC_KEY_FILE", "server/certs/quic-gateway.key")) alpn = [item.strip() for item in os.getenv("QUIC_ALPN", "cmvr-quic-edge/1").split(",") if item.strip()] configuration = QuicConfiguration(is_client=False, alpn_protocols=alpn, max_datagram_frame_size=65536) configuration.load_cert_chain(cert, key) loop = asyncio.get_running_loop() _, server = await loop.create_datagram_endpoint( lambda: LoggingQuicServer( configuration=configuration, create_protocol=CmvrQuicProtocol, ), local_addr=(host, port), ) log("listening", address=f"{host}:{port}", alpn=alpn, cert=str(cert)) stopped = asyncio.Event() loop = asyncio.get_running_loop() for sig in (signal.SIGINT, signal.SIGTERM): try: loop.add_signal_handler(sig, stopped.set) except NotImplementedError: signal.signal(sig, lambda *_: loop.call_soon_threadsafe(stopped.set)) await stopped.wait() server.close() if __name__ == "__main__": asyncio.run(main())