diff options
Diffstat (limited to 'packages/meshbay-node/src/meshbay_node/transport/quic_server.py')
| -rw-r--r-- | packages/meshbay-node/src/meshbay_node/transport/quic_server.py | 176 |
1 files changed, 167 insertions, 9 deletions
diff --git a/packages/meshbay-node/src/meshbay_node/transport/quic_server.py b/packages/meshbay-node/src/meshbay_node/transport/quic_server.py index 9cb3bd8..43c1026 100644 --- a/packages/meshbay-node/src/meshbay_node/transport/quic_server.py +++ b/packages/meshbay-node/src/meshbay_node/transport/quic_server.py @@ -21,6 +21,7 @@ import asyncio import base64 import logging import struct +import subprocess from pathlib import Path from typing import Any, Callable @@ -50,6 +51,25 @@ MAX_MSG = 64 * 1024 * 1024 ALPN = ["meshbay-mnp"] +class Denylist: + """Shared denylist for revoked users and invalidated JWTs.""" + + def __init__(self): + self.user_ids: set[str] = set() + self.jtis: set[str] = set() + + def is_denied(self, user_id: str, jti: str) -> bool: + return user_id in self.user_ids or jti in self.jtis + + def deny_user(self, user_id: str) -> None: + self.user_ids.add(user_id) + log.info("Denied user: %s", user_id[:8]) + + def deny_jti(self, jti: str) -> None: + self.jtis.add(jti) + log.info("Denied jti: %s", jti[:8]) + + # ── Wire helpers ────────────────────────────────────────────────────────────── def _pack(obj: dict) -> bytes: @@ -90,6 +110,7 @@ class _MNPServerProtocol(QuicConnectionProtocol): super().__init__(*args, **kwargs) self._ctx = node_ctx # shared server context (keys, index, etc.) self._user_id: str | None = None + self._group_id: str | None = None self._buffers: dict[int, _StreamBuffer] = {} def quic_event_received(self, event: QuicEvent) -> None: @@ -116,6 +137,10 @@ class _MNPServerProtocol(QuicConnectionProtocol): self._do_index_sync_sync(stream_id) elif mtype == MNP.FILE_REQUEST: self._do_file_request_sync(stream_id, msg) + elif mtype == MNP.STREAM_SEGMENT: + self._do_stream_segment_sync(stream_id, msg) + elif mtype == MNP.CHAT_MESSAGE: + self._do_chat_message_sync(stream_id, msg) else: log.warning("Unknown MNP message type: %s", mtype) except Exception as e: @@ -123,8 +148,8 @@ class _MNPServerProtocol(QuicConnectionProtocol): 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", "") + group_id = msg.get("group_id", "") try: decoded = jwt.decode(token, self._ctx["hub_pk_pem"], algorithms=["EdDSA"]) except Exception as e: @@ -132,21 +157,45 @@ class _MNPServerProtocol(QuicConnectionProtocol): self._quic.close() return - if decoded.get("exp", 0) < int(time.time()): - self._send(stream_id, {"type": "error", "detail": "JWT expired"}) + denylist = self._ctx.get("denylist") + if denylist and denylist.is_denied(decoded.get("sub", ""), decoded.get("jti", "")): + self._send(stream_id, {"type": "error", "detail": "Token revoked"}) + self._quic.close() + return + + if group_id and group_id not in decoded.get("groups", []): + self._send(stream_id, {"type": "error", "detail": "Not a member of this group"}) + self._quic.close() + return + + if group_id and "groups" in self._ctx and group_id not in self._ctx["groups"]: + self._send(stream_id, {"type": "error", "detail": "Group not hosted on this node"}) self._quic.close() return self._user_id = decoded["sub"] - log.info("QUIC handshake OK — user=%s", self._user_id[:8]) + self._group_id = group_id + + peers = self._ctx.get("_peers") + if peers is not None: + peers[self._user_id] = self + + log.info("QUIC handshake OK — user=%s group=%s", self._user_id[:8], group_id[:8] if group_id else "none") self._send(stream_id, { "type": MNP.HANDSHAKE_ACK, "v": MNP_VERSION, "node_pk": pk_to_b64(self._ctx["sk_node"].public_key()), }) + def _group_ctx(self) -> dict: + """Resolve the active group context (multi-group or legacy single-group).""" + if "groups" in self._ctx and self._group_id: + return self._ctx["groups"][self._group_id] + return self._ctx + def _do_index_sync_sync(self, stream_id: int) -> None: - wire = self._ctx["index"].serialize() + ctx = self._group_ctx() + wire = ctx["index"].serialize() self._send(stream_id, { "type": MNP.INDEX_SYNC, "v": MNP_VERSION, @@ -155,26 +204,96 @@ class _MNPServerProtocol(QuicConnectionProtocol): def _do_file_request_sync(self, stream_id: int, msg: dict) -> None: """Serve file chunk synchronously (blocking I/O — acceptable for test sizes).""" + ctx = self._group_ctx() file_id = msg["file_id"] chunk_index = msg["chunk_index"] - entry = self._ctx["index"].get_entry(file_id) + entry = 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 + file_path = 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"], + ctx["gek"], file_path, chunk_index, ) self._send(stream_id, chunk_data) + def _do_stream_segment_sync(self, stream_id: int, msg: dict) -> None: + """Extract and serve one HLS segment via ffmpeg.""" + 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(stream_id, {"type": "error", "detail": "File not found"}) + return + + file_path = ctx["shared_root"] / entry.path / entry.name + if not file_path.exists(): + self._send(stream_id, {"type": "error", "detail": "File not on disk"}) + return + + start_time = segment_index * segment_duration + segment_data = _extract_segment(file_path, start_time, segment_duration) + if segment_data is None: + self._send(stream_id, {"type": "error", "detail": "Segment extraction failed"}) + return + + self._send(stream_id, { + "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_sync(self, stream_id: int, msg: dict) -> None: + """Receive a chat message, store it, and broadcast to other connected peers.""" + chat_store = self._ctx.get("chat_store") + if chat_store: + import asyncio + 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"), + )) + + peers = self._ctx.get("_peers", {}) + broadcast = { + "type": MNP.CHAT_MESSAGE, + "v": MNP_VERSION, + "sender_id": msg.get("sender_id", self._user_id), + "iteration": msg.get("iteration", 0), + "payload": msg.get("payload", ""), + "thread_id": msg.get("thread_id"), + "group_id": self._group_id or "", + } + for uid, proto in peers.items(): + if uid != self._user_id and proto is not self: + try: + proto._send(0, broadcast) + except Exception: + pass + + self._send(stream_id, {"type": "ack", "v": MNP_VERSION}) + + def connection_lost(self, exc) -> None: + peers = self._ctx.get("_peers") + if peers and self._user_id: + peers.pop(self._user_id, None) + super().connection_lost(exc) + def _send(self, stream_id: int, obj: dict) -> None: self._quic.send_stream_data(stream_id, _pack(obj)) self.transmit() @@ -213,6 +332,25 @@ def _read_and_encrypt( } +def _extract_segment(file_path: Path, start_time: float, duration: float) -> bytes | None: + """Extract one HLS segment via ffmpeg. Returns MPEG-TS bytes or None on failure.""" + try: + result = subprocess.run( + ["ffmpeg", "-hide_banner", "-loglevel", "error", + "-ss", str(start_time), + "-i", str(file_path), + "-t", str(duration), + "-c:v", "copy", "-c:a", "copy", + "-f", "mpegts", "pipe:1"], + capture_output=True, timeout=30, + ) + if result.returncode == 0 and result.stdout: + return result.stdout + return None + except Exception: + return None + + # ── QuicChunkServer ──────────────────────────────────────────────────────────── class QuicChunkServer: @@ -228,10 +366,12 @@ class QuicChunkServer: gek: bytes, shared_root: Path, index: GroupIndex, - host: str = "::", # écoute IPv4 + IPv6 (dual-stack Linux) + host: str = "::", # listen IPv4 + IPv6 (dual-stack Linux) port: int = 19000, cert_path: Path | None = None, key_path: Path | None = None, + groups: dict[str, dict] | None = None, + denylist: Denylist | None = None, ): self._ctx = { "sk_node": sk_node, @@ -240,17 +380,27 @@ class QuicChunkServer: "shared_root": shared_root, "index": index, } + if groups: + self._ctx["groups"] = groups + self._denylist = denylist or Denylist() + self._ctx["denylist"] = self._denylist + self._ctx["_peers"] = {} 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 + self._session_tickets: dict[bytes, Any] = {} @property def port(self) -> int: return self._port + @property + def denylist(self) -> Denylist: + return self._denylist + def _make_config(self) -> QuicConfiguration: from meshbay_node.transport.tls_cert import generate_self_signed_cert if not self._cert_path.exists(): @@ -259,6 +409,12 @@ class QuicChunkServer: config.load_cert_chain(str(self._cert_path), str(self._key_path)) return config + def _store_ticket(self, ticket: Any) -> None: + self._session_tickets[ticket.ticket] = ticket + + def _fetch_ticket(self, label: bytes) -> Any | None: + return self._session_tickets.pop(label, None) + async def start(self) -> None: config = self._make_config() ctx = self._ctx @@ -270,6 +426,8 @@ class QuicChunkServer: self._host, self._port, configuration=config, create_protocol=protocol_factory, + session_ticket_handler=self._store_ticket, + session_ticket_fetcher=self._fetch_ticket, ) log.info("QuicChunkServer listening on %s:%d (QUIC/UDP)", self._host, self._port) |