259 lines
11 KiB
Python
259 lines
11 KiB
Python
|
|
"""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())
|