diff options
Diffstat (limited to 'packages/meshbay-node/src')
3 files changed, 399 insertions, 3 deletions
diff --git a/packages/meshbay-node/src/meshbay_node/hub_client.py b/packages/meshbay-node/src/meshbay_node/hub_client.py index 74851c1..3e77ed0 100644 --- a/packages/meshbay-node/src/meshbay_node/hub_client.py +++ b/packages/meshbay-node/src/meshbay_node/hub_client.py @@ -250,10 +250,11 @@ class HubClient: self, on_incoming: Any = None, on_revocation: Any = None, + on_webrtc_offer: Any = None, ) -> None: """ Maintain a persistent WebSocket connection to the hub. - Receives NAT punch requests and revocation tokens. + Receives NAT punch requests, revocation tokens, and WebRTC offers. Runs until cancelled. """ import websockets @@ -270,6 +271,7 @@ class HubClient: await ws.send(json.dumps({ "type": "auth", "token": self._session.access_token, + "node_id": self._session.node_id, })) auth_resp = json.loads(await ws.recv()) if auth_resp.get("type") != "auth_ok": @@ -289,6 +291,18 @@ class HubClient: elif mtype == "revocation" and on_revocation: on_revocation(msg.get("token", "")) + elif mtype == "webrtc_offer" and on_webrtc_offer: + answer = await on_webrtc_offer( + msg["sdp"], msg["peer_id"], + msg.get("ice_candidates", [])) + if answer: + await ws.send(json.dumps({ + "type": "webrtc_answer", + "peer_id": msg["peer_id"], + "sdp": answer[0], + "ice_candidates": answer[1], + })) + elif mtype == "pong": pass diff --git a/packages/meshbay-node/src/meshbay_node/transport/__init__.py b/packages/meshbay-node/src/meshbay_node/transport/__init__.py index 5a1b8d7..df9c209 100644 --- a/packages/meshbay-node/src/meshbay_node/transport/__init__.py +++ b/packages/meshbay-node/src/meshbay_node/transport/__init__.py @@ -1,10 +1,9 @@ -"""MeshBay Node transport layer — TCP+TLS (MNP v1) and QUIC (MNP v2).""" +"""MeshBay Node transport layer — TCP+TLS (v1), QUIC (v2), WebRTC (browsers).""" from .server import ChunkServer from .client import ChunkClient from .http_server import create_http_app # QUIC transport (MNP v2) — requires aioquic>=1.0 -# Falls back gracefully if not installed; node still works via TCP+TLS and HTTP. try: from .quic_server import QuicChunkServer, Denylist from .quic_client import QuicChunkClient @@ -15,7 +14,17 @@ except ImportError: Denylist = None # type: ignore[assignment,misc] QUIC_AVAILABLE = False +# WebRTC transport (browsers) — requires aiortc>=1.9 +try: + from .webrtc_server import WebRTCTransport, WebRTCPeerSession + WEBRTC_AVAILABLE = True +except ImportError: + WebRTCTransport = None # type: ignore[assignment,misc] + WebRTCPeerSession = None # type: ignore[assignment,misc] + WEBRTC_AVAILABLE = False + __all__ = [ "ChunkServer", "ChunkClient", "create_http_app", "QuicChunkServer", "QuicChunkClient", "Denylist", "QUIC_AVAILABLE", + "WebRTCTransport", "WebRTCPeerSession", "WEBRTC_AVAILABLE", ] diff --git a/packages/meshbay-node/src/meshbay_node/transport/webrtc_server.py b/packages/meshbay-node/src/meshbay_node/transport/webrtc_server.py new file mode 100644 index 0000000..89391b9 --- /dev/null +++ b/packages/meshbay-node/src/meshbay_node/transport/webrtc_server.py @@ -0,0 +1,373 @@ +""" +MeshBay Node — WebRTC DataChannel server for browser clients. + +Browsers cannot use QUIC for NAT traversal (WebTransport doesn't allow choosing +the UDP source port — Port-Restricted Cone NAT requires exact port matching). +WebRTC DataChannel with ICE/STUN handles this automatically. + +The MNP protocol (handshake, file_request, file_chunk, chat, etc.) runs +identically over WebRTC DataChannel as over QUIC streams. Same E2E encryption, +same message types, same msgpack wire format. + +Wire format on the DataChannel: + - Each message is length-prefixed msgpack (4-byte big-endian + msgpack payload) + - Same as QUIC streams and TCP+TLS + - DataChannel is ordered and reliable (SCTP over DTLS) + +Signaling flow (handled externally by the hub): + Browser → Hub : POST /v1/nodes/{id}/webrtc/offer {sdp, ice_candidates} + Hub → Node : WS push {type: "webrtc_offer", sdp, ice_candidates, peer_id} + Node → Hub : WS push {type: "webrtc_answer", sdp, ice_candidates, peer_id} + Hub → Browser : SSE/response {sdp, ice_candidates} + After signaling, DataChannel is P2P — hub is out of the loop. +""" + +import asyncio +import base64 +import logging +import struct +from pathlib import Path +from typing import Any + +import blake3 +import jwt +import msgpack +from aiortc import RTCPeerConnection, RTCSessionDescription, RTCDataChannel +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 + +log = logging.getLogger(__name__) + +CHUNK_SIZE = 1024 * 1024 +MAX_MSG = 64 * 1024 * 1024 + + +def _pack(obj: dict) -> bytes: + data = msgpack.packb(obj, use_bin_type=True) + return struct.pack(">I", len(data)) + data + + +class _DataChannelBuffer: + """Accumulate DataChannel messages and extract length-prefixed msgpack.""" + + 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) + + +class WebRTCPeerSession: + """One WebRTC peer connection, handling MNP over a DataChannel.""" + + def __init__(self, pc: RTCPeerConnection, node_ctx: dict): + self._pc = pc + self._ctx = node_ctx + self._channel: RTCDataChannel | None = None + self._buffer = _DataChannelBuffer() + self._user_id: str | None = None + self._group_id: str | None = None + + def _setup_channel(self, channel: RTCDataChannel) -> None: + self._channel = channel + + @channel.on("message") + def on_message(message): + if isinstance(message, str): + message = message.encode() + self._buffer.feed(message) + for msg in self._buffer.messages(): + self._handle_message(msg) + + def _handle_message(self, msg: dict) -> None: + mtype = msg.get("type") + try: + if mtype == MNP.HANDSHAKE: + self._do_handshake(msg) + elif self._user_id is None: + self._send({"type": "error", "detail": "Handshake required"}) + elif mtype == MNP.INDEX_SYNC: + self._do_index_sync() + elif mtype == MNP.FILE_REQUEST: + self._do_file_request(msg) + elif mtype == MNP.STREAM_SEGMENT: + self._do_stream_segment(msg) + elif mtype == MNP.CHAT_MESSAGE: + self._do_chat_message(msg) + else: + log.warning("Unknown MNP message type on DataChannel: %s", mtype) + except Exception as e: + log.error("Error handling %s on DataChannel: %s", mtype, e) + self._send({"type": "error", "detail": str(e)}) + + def _do_handshake(self, msg: dict) -> None: + token = msg.get("token", "") + group_id = msg.get("group_id", "") + try: + decoded = jwt.decode(token, self._ctx["hub_pk_pem"], algorithms=["EdDSA"]) + except Exception as e: + self._send({"type": "error", "detail": f"Invalid JWT: {e}"}) + return + + denylist = self._ctx.get("denylist") + if denylist and denylist.is_denied(decoded.get("sub", ""), decoded.get("jti", "")): + self._send({"type": "error", "detail": "Token revoked"}) + return + + if group_id and group_id not in decoded.get("groups", []): + self._send({"type": "error", "detail": "Not a member of this group"}) + return + + if group_id and "groups" in self._ctx and group_id not in self._ctx["groups"]: + self._send({"type": "error", "detail": "Group not hosted on this node"}) + return + + self._user_id = decoded["sub"] + self._group_id = group_id + + log.info("WebRTC handshake OK — user=%s group=%s", + self._user_id[:8], group_id[:8] if group_id else "none") + self._send({ + "type": MNP.HANDSHAKE_ACK, + "v": MNP_VERSION, + "node_pk": pk_to_b64(self._ctx["sk_node"].public_key()), + }) + + def _group_ctx(self) -> dict: + if "groups" in self._ctx and self._group_id: + return self._ctx["groups"][self._group_id] + return self._ctx + + def _do_index_sync(self) -> None: + ctx = self._group_ctx() + wire = ctx["index"].serialize() + self._send({ + "type": MNP.INDEX_SYNC, + "v": MNP_VERSION, + "index_b64": base64.b64encode(wire).decode(), + }) + + def _do_file_request(self, msg: dict) -> None: + ctx = self._group_ctx() + file_id = msg["file_id"] + chunk_index = msg["chunk_index"] + entry = ctx["index"].get_entry(file_id) + if not entry: + self._send({"type": "error", "detail": "File not found"}) + return + + file_path = ctx["shared_root"] / entry.path / entry.name + if not file_path.exists(): + self._send({"type": "error", "detail": "File not on disk"}) + return + + chunk_data = _read_and_encrypt( + self._ctx["sk_node"], + ctx["gek"], + file_path, + chunk_index, + ) + self._send(chunk_data) + + def _do_stream_segment(self, msg: dict) -> None: + ctx = self._group_ctx() + file_id = msg["file_id"] + segment_index = msg["segment_index"] + segment_duration = msg.get("segment_duration", 4) + + entry = ctx["index"].get_entry(file_id) + if not entry: + self._send({"type": "error", "detail": "File not found"}) + return + + file_path = ctx["shared_root"] / entry.path / entry.name + if not file_path.exists(): + self._send({"type": "error", "detail": "File not on disk"}) + return + + import subprocess + try: + result = subprocess.run( + ["ffmpeg", "-hide_banner", "-loglevel", "error", + "-ss", str(segment_index * segment_duration), + "-i", str(file_path), + "-t", str(segment_duration), + "-c:v", "copy", "-c:a", "copy", + "-f", "mpegts", "pipe:1"], + capture_output=True, timeout=30, + ) + if result.returncode != 0 or not result.stdout: + self._send({"type": "error", "detail": "Segment extraction failed"}) + return + segment_data = result.stdout + except Exception: + self._send({"type": "error", "detail": "Segment extraction failed"}) + return + + self._send({ + "type": MNP.STREAM_SEGMENT, + "v": MNP_VERSION, + "file_id": file_id, + "segment_index": segment_index, + "data_b64": base64.b64encode(segment_data).decode(), + "size": len(segment_data), + }) + + def _do_chat_message(self, msg: dict) -> None: + chat_store = self._ctx.get("chat_store") + if chat_store: + asyncio.ensure_future(chat_store.save_message( + sender_id=msg.get("sender_id", self._user_id), + iteration=msg.get("iteration", 0), + payload=msg.get("payload", b"").encode() + if isinstance(msg.get("payload"), str) else msg.get("payload", b""), + thread_id=msg.get("thread_id"), + )) + self._send({"type": "ack", "v": MNP_VERSION}) + + def _send(self, obj: dict) -> None: + if self._channel and self._channel.readyState == "open": + self._channel.send(_pack(obj)) + + async def close(self) -> None: + await self._pc.close() + + +def _read_and_encrypt( + sk_node: Ed25519PrivateKey, + gek: bytes, + file_path: Path, + chunk_index: int, +) -> dict: + 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(), + } + + +class WebRTCTransport: + """ + Manages WebRTC peer connections for browser clients. + + Usage: + transport = WebRTCTransport(sk_node, hub_pk_pem, gek, shared_root, index) + answer_sdp = await transport.handle_offer(offer_sdp, peer_id) + # Return answer_sdp to the browser via hub signaling + """ + + def __init__( + self, + sk_node: Ed25519PrivateKey, + hub_pk_pem: bytes, + gek: bytes, + shared_root: Path, + index: GroupIndex, + groups: dict[str, dict] | None = None, + denylist: Any | None = None, + stun_servers: list[str] | None = None, + ): + self._ctx: dict[str, Any] = { + "sk_node": sk_node, + "hub_pk_pem": hub_pk_pem, + "gek": gek, + "shared_root": shared_root, + "index": index, + } + if groups: + self._ctx["groups"] = groups + if denylist: + self._ctx["denylist"] = denylist + self._stun = stun_servers or ["stun:stun.l.google.com:19302"] + self._sessions: dict[str, WebRTCPeerSession] = {} + + async def handle_offer( + self, offer_sdp: str, peer_id: str, + ) -> tuple[str, list[dict]]: + """ + Process a WebRTC SDP offer from a browser client. + + Returns (answer_sdp, ice_candidates) to relay back via hub signaling. + ICE candidates are embedded in the SDP (aiortc gathers before returning). + """ + from aiortc import RTCIceServer, RTCConfiguration + + config = RTCConfiguration( + iceServers=[RTCIceServer(urls=s) for s in self._stun] if self._stun else [] + ) + pc = RTCPeerConnection(configuration=config) + session = WebRTCPeerSession(pc, self._ctx) + self._sessions[peer_id] = session + + @pc.on("datachannel") + def on_datachannel(channel: RTCDataChannel): + log.info("WebRTC DataChannel opened: %s (peer=%s)", channel.label, peer_id) + session._setup_channel(channel) + + @pc.on("connectionstatechange") + async def on_state_change(): + state = pc.connectionState + log.info("WebRTC connection state: %s (peer=%s)", state, peer_id) + if state in ("failed", "closed"): + self._sessions.pop(peer_id, None) + + offer = RTCSessionDescription(sdp=offer_sdp, type="offer") + await pc.setRemoteDescription(offer) + answer = await pc.createAnswer() + await pc.setLocalDescription(answer) + + log.info("WebRTC answer ready for peer=%s", peer_id) + return pc.localDescription.sdp, [] + + async def close_peer(self, peer_id: str) -> None: + session = self._sessions.pop(peer_id, None) + if session: + await session.close() + + async def close_all(self) -> None: + for session in self._sessions.values(): + await session.close() + self._sessions.clear() + + @property + def active_peers(self) -> int: + return len(self._sessions) |