inspection-host-computer/server/quic/server.py

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())