""" MeshBay Node — QUIC chunk server (MNP v2). Replaces the TCP+TLS ChunkServer with QUIC transport. Advantages over TCP: - UDP-based → works with hole punching (Spike 4 confirmed Cone NAT on SFR) - Multiplexed streams — each request is an independent QUIC stream - 0-RTT reconnection (connection resumption) - Built-in TLS 1.3 Wire protocol: - Each bidirectional QUIC stream carries one request/response exchange - Messages: length-prefixed msgpack (4-byte big-endian, same as TCP+TLS) - MNP handshake on stream 0 (control stream); subsequent streams = requests Application protocol (MNP) is identical to TCP+TLS version. The transport is the only change — all crypto, auth, and message types stay the same. """ import asyncio import base64 import logging import struct from pathlib import Path from typing import Any, Callable import blake3 import jwt import msgpack from aioquic.asyncio import QuicConnectionProtocol, serve from aioquic.quic.configuration import QuicConfiguration from aioquic.quic.events import QuicEvent, StreamDataReceived, StreamReset from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey from meshbay_common import MNP_VERSION from meshbay_common.crypto import ( chunk_key as derive_chunk_key, encrypt_chunk, sign_chunk, pk_to_b64, ) from meshbay_common.protocol import MNP from meshbay_node.indexer import GroupIndex from meshbay_node.transport.tls_cert import server_ssl_context log = logging.getLogger(__name__) CHUNK_SIZE = 1024 * 1024 MAX_MSG = 64 * 1024 * 1024 ALPN = ["meshbay-mnp"] # ── Wire helpers ────────────────────────────────────────────────────────────── def _pack(obj: dict) -> bytes: data = msgpack.packb(obj, use_bin_type=True) return struct.pack(">I", len(data)) + data class _StreamBuffer: """Accumulate incoming QUIC stream data and extract length-prefixed messages.""" def __init__(self): self._buf = bytearray() def feed(self, data: bytes): self._buf.extend(data) def messages(self): while len(self._buf) >= 4: length = struct.unpack(">I", self._buf[:4])[0] if length > MAX_MSG: raise ValueError(f"Message too large: {length}") if len(self._buf) < 4 + length: break msg_bytes = bytes(self._buf[4:4 + length]) del self._buf[:4 + length] yield msgpack.unpackb(msg_bytes, raw=False) # ── Per-connection server protocol ──────────────────────────────────────────── class _MNPServerProtocol(QuicConnectionProtocol): """ One instance per QUIC connection. Handles the MNP handshake and all subsequent streams. """ def __init__(self, *args, node_ctx: dict, **kwargs): super().__init__(*args, **kwargs) self._ctx = node_ctx # shared server context (keys, index, etc.) self._user_id: str | None = None self._buffers: dict[int, _StreamBuffer] = {} def quic_event_received(self, event: QuicEvent) -> None: if isinstance(event, StreamDataReceived): sid = event.stream_id if sid not in self._buffers: self._buffers[sid] = _StreamBuffer() self._buffers[sid].feed(event.data) for msg in self._buffers[sid].messages(): self._handle_message_sync(sid, msg) elif isinstance(event, StreamReset): self._buffers.pop(event.stream_id, None) def _handle_message_sync(self, stream_id: int, msg: dict) -> None: """Handle an MNP message synchronously (called from quic_event_received).""" mtype = msg.get("type") try: if mtype == MNP.HANDSHAKE: self._do_handshake_sync(stream_id, msg) elif self._user_id is None: self._send(stream_id, {"type": "error", "detail": "Handshake required"}) elif mtype == MNP.INDEX_SYNC: self._do_index_sync_sync(stream_id) elif mtype == MNP.FILE_REQUEST: self._do_file_request_sync(stream_id, msg) else: log.warning("Unknown MNP message type: %s", mtype) except Exception as e: log.error("Error handling %s: %s", mtype, e) self._send(stream_id, {"type": "error", "detail": str(e)}) def _do_handshake_sync(self, stream_id: int, msg: dict) -> None: import time token = msg.get("token", "") try: decoded = jwt.decode(token, self._ctx["hub_pk_pem"], algorithms=["EdDSA"]) except Exception as e: self._send(stream_id, {"type": "error", "detail": f"Invalid JWT: {e}"}) self._quic.close() return if decoded.get("exp", 0) < int(time.time()): self._send(stream_id, {"type": "error", "detail": "JWT expired"}) self._quic.close() return self._user_id = decoded["sub"] log.info("QUIC handshake OK — user=%s", self._user_id[:8]) self._send(stream_id, { "type": MNP.HANDSHAKE_ACK, "v": MNP_VERSION, "node_pk": pk_to_b64(self._ctx["sk_node"].public_key()), }) def _do_index_sync_sync(self, stream_id: int) -> None: wire = self._ctx["index"].serialize() self._send(stream_id, { "type": MNP.INDEX_SYNC, "v": MNP_VERSION, "index_b64": base64.b64encode(wire).decode(), }) def _do_file_request_sync(self, stream_id: int, msg: dict) -> None: """Serve file chunk synchronously (blocking I/O — acceptable for test sizes).""" file_id = msg["file_id"] chunk_index = msg["chunk_index"] entry = self._ctx["index"].get_entry(file_id) if not entry: self._send(stream_id, {"type": "error", "detail": "File not found"}) return file_path = self._ctx["shared_root"] / entry.path / entry.name if not file_path.exists(): self._send(stream_id, {"type": "error", "detail": "File not on disk"}) return chunk_data = _read_and_encrypt( self._ctx["sk_node"], self._ctx["gek"], file_path, chunk_index, ) self._send(stream_id, chunk_data) def _send(self, stream_id: int, obj: dict) -> None: self._quic.send_stream_data(stream_id, _pack(obj)) self.transmit() def _read_and_encrypt( sk_node: Ed25519PrivateKey, gek: bytes, file_path: Path, chunk_index: int, ) -> dict: """Read and encrypt one chunk (blocking — runs in executor).""" with open(file_path, "rb") as f: f.seek(chunk_index * CHUNK_SIZE) plaintext = f.read(CHUNK_SIZE) file_hash = blake3.blake3(file_path.read_bytes()).digest() pt_hash = blake3.blake3(plaintext).digest() ckey = derive_chunk_key(gek, file_hash, chunk_index) nonce, ct = encrypt_chunk(ckey, plaintext) ct_hash = blake3.blake3(ct).digest() sig = sign_chunk(sk_node, chunk_index, nonce, ct_hash) return { "type": MNP.FILE_CHUNK, "v": MNP_VERSION, "chunk_index": chunk_index, "plaintext_size": len(plaintext), "nonce_b64": base64.b64encode(nonce).decode(), "ct_b64": base64.b64encode(ct).decode(), "ct_hash_b64": base64.b64encode(ct_hash).decode(), "pt_hash_b64": base64.b64encode(pt_hash).decode(), "sig_b64": base64.b64encode(sig).decode(), "pk_node_b64": pk_to_b64(sk_node.public_key()), "file_hash_b64": base64.b64encode(file_hash).decode(), } # ── QuicChunkServer ──────────────────────────────────────────────────────────── class QuicChunkServer: """ QUIC-based MNP chunk server (MNP v2). Drop-in replacement for ChunkServer with UDP transport. """ def __init__( self, sk_node: Ed25519PrivateKey, hub_pk_pem: bytes, gek: bytes, shared_root: Path, index: GroupIndex, host: str = "0.0.0.0", port: int = 19000, cert_path: Path | None = None, key_path: Path | None = None, ): self._ctx = { "sk_node": sk_node, "hub_pk_pem": hub_pk_pem, "gek": gek, "shared_root": shared_root, "index": index, } self._host = host self._port = port self._cert_path = cert_path or Path.home() / ".config/meshbay/node_tls.crt" self._key_path = key_path or Path.home() / ".config/meshbay/node_tls.key" self._server = None self._task = None @property def port(self) -> int: return self._port def _make_config(self) -> QuicConfiguration: from meshbay_node.transport.tls_cert import generate_self_signed_cert if not self._cert_path.exists(): generate_self_signed_cert(self._cert_path, self._key_path) config = QuicConfiguration(is_client=False, alpn_protocols=ALPN) config.load_cert_chain(str(self._cert_path), str(self._key_path)) return config async def start(self) -> None: config = self._make_config() ctx = self._ctx def protocol_factory(*args, **kwargs): return _MNPServerProtocol(*args, node_ctx=ctx, **kwargs) self._server = await serve( self._host, self._port, configuration=config, create_protocol=protocol_factory, ) log.info("QuicChunkServer listening on %s:%d (QUIC/UDP)", self._host, self._port) async def stop(self) -> None: if self._server: self._server.close() self._server = None log.info("QuicChunkServer stopped")