""" MeshBay Node — TCP+TLS chunk server (MNP v1). Serves encrypted file chunks to authenticated clients over TLS. Each connection: 1. Client sends MNP handshake with JWT bearer token 2. Server verifies JWT offline (hub PK cached) 3. Client sends chunk requests 4. Server reads from disk, encrypts on-the-fly, signs, sends Wire protocol: length-prefixed msgpack (4-byte big-endian length header). All messages carry {"type": ..., "v": MNP_VERSION}. """ import asyncio import base64 import logging import struct import time from pathlib import Path import blake3 import jwt import msgpack 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 # 1 MB MAX_MSG = 64 * 1024 * 1024 # 64 MB max message size (safety) # ── Wire helpers ────────────────────────────────────────────────────────────── async def _send(writer: asyncio.StreamWriter, obj: dict) -> None: data = msgpack.packb(obj, use_bin_type=True) writer.write(struct.pack(">I", len(data)) + data) await writer.drain() async def _recv(reader: asyncio.StreamReader) -> dict: header = await reader.readexactly(4) length = struct.unpack(">I", header)[0] if length > MAX_MSG: raise ValueError(f"Message too large: {length}") data = await reader.readexactly(length) return msgpack.unpackb(data, raw=False) # ── Chunk serving ───────────────────────────────────────────────────────────── def _serve_chunk( sk_node: Ed25519PrivateKey, gek: bytes, file_path: Path, file_hash: bytes, chunk_index: int, ) -> dict: """Read, encrypt, sign one chunk. Blocking — run in executor.""" with open(file_path, "rb") as f: f.seek(chunk_index * CHUNK_SIZE) plaintext = f.read(CHUNK_SIZE) 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(), } # ── Connection handler ──────────────────────────────────────────────────────── class _ConnectionHandler: def __init__( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter, sk_node: Ed25519PrivateKey, hub_pk_pem: bytes, gek: bytes, shared_root: Path, index: GroupIndex, ): self._reader = reader self._writer = writer self._sk_node = sk_node self._hub_pk_pem = hub_pk_pem self._gek = gek self._shared_root = shared_root self._index = index self._peer = writer.get_extra_info("peername") self._user_id: str | None = None async def handle(self) -> None: try: await self._handshake() await self._serve_loop() except asyncio.IncompleteReadError: log.debug("[%s] Client disconnected", self._peer) except Exception as e: log.warning("[%s] Error: %s", self._peer, e) await _send(self._writer, {"type": "error", "detail": str(e)}) finally: self._writer.close() async def _handshake(self) -> None: msg = await _recv(self._reader) if msg.get("type") != MNP.HANDSHAKE: raise ValueError(f"Expected handshake, got {msg.get('type')!r}") token = msg.get("token", "") try: decoded = jwt.decode(token, self._hub_pk_pem, algorithms=["EdDSA"]) except Exception as e: raise PermissionError(f"Invalid JWT: {e}") from e if decoded.get("exp", 0) < int(time.time()): raise PermissionError("JWT expired") self._user_id = decoded["sub"] log.info("[%s] Handshake OK — user=%s", self._peer, self._user_id[:8]) await _send(self._writer, { "type": MNP.HANDSHAKE_ACK, "v": MNP_VERSION, "node_pk": pk_to_b64(self._sk_node.public_key()), }) async def _serve_loop(self) -> None: loop = asyncio.get_event_loop() while True: msg = await _recv(self._reader) mtype = msg.get("type") if mtype == MNP.INDEX_SYNC: wire = self._index.serialize() await _send(self._writer, { "type": MNP.INDEX_SYNC, "v": MNP_VERSION, "index_b64": base64.b64encode(wire).decode(), }) elif mtype == MNP.FILE_REQUEST: file_id = msg["file_id"] chunk_index = msg["chunk_index"] entry = self._index.get_entry(file_id) if entry is None: await _send(self._writer, { "type": "error", "detail": f"File not found: {file_id[:8]}", }) continue file_path = self._shared_root / entry.path / entry.name if not file_path.exists(): await _send(self._writer, { "type": "error", "detail": "File not on disk"}) continue file_hash = blake3.blake3(file_path.read_bytes()).digest() chunk = await loop.run_in_executor( None, _serve_chunk, self._sk_node, self._gek, file_path, file_hash, chunk_index) await _send(self._writer, chunk) else: log.warning("[%s] Unknown message type: %s", self._peer, mtype) # ── Server ──────────────────────────────────────────────────────────────────── class ChunkServer: """ Async TCP+TLS server that serves encrypted file chunks. Usage: server = ChunkServer( host="0.0.0.0", port=19000, sk_node=sk, hub_pk_pem=pk_pem, gek=gek, shared_root=Path("/data"), index=group_index, ) await server.start() # ... when shutting down: await server.stop() """ 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._sk_node = sk_node self._hub_pk_pem = hub_pk_pem self._gek = gek self._shared_root = shared_root self._index = index self._host = host self._port = port self._cert_path = cert_path self._key_path = key_path self._server: asyncio.Server | None = None @property def port(self) -> int: return self._port async def start(self) -> None: ssl_ctx = server_ssl_context( cert_path=self._cert_path or Path.home() / ".config/meshbay/node_tls.crt", key_path=self._key_path or Path.home() / ".config/meshbay/node_tls.key", ) self._server = await asyncio.start_server( self._handle_connection, host=self._host, port=self._port, ssl=ssl_ctx, ) log.info("ChunkServer listening on %s:%d (TLS)", self._host, self._port) async def stop(self) -> None: if self._server: self._server.close() await self._server.wait_closed() self._server = None log.info("ChunkServer stopped") async def _handle_connection( self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter ) -> None: handler = _ConnectionHandler( reader, writer, self._sk_node, self._hub_pk_pem, self._gek, self._shared_root, self._index, ) await handler.handle()