"""Collaboration temps réel — édition simultanée (ROADMAP #62). Ce module implémente le cœur serveur de l'édition collaborative : * un **relais WebSocket** : une *room* est créée par fichier ouvert (clé ``vault::chemin``) et tous les clients qui éditent le même fichier rejoignent la même room ; * un **relais de mises à jour Yjs** (CRDT) : le serveur ne décode pas le format binaire Yjs, il stocke le journal des mises à jour reçues et le rejoue aux nouveaux arrivants. La fusion sans conflit est assurée côté client par Yjs ; * un **awareness** (curseurs colorés + sélections) relayé entre clients ; * une **persistance différée** : le texte markdown reçu des clients est écrit sur disque après un debounce (2 s par défaut). Sécurité : chaque connexion est authentifiée manuellement (les dépendances FastAPI ``Depends`` ne s'exécutent pas pour ``@app.websocket``), puis le chemin est validé via :func:`backend.services.paths.resolve_safe_path` et l'accès à la vault via ``check_vault_access``. """ from __future__ import annotations import asyncio import base64 import json import logging import time from dataclasses import dataclass, field from pathlib import Path from typing import Any from fastapi import WebSocket from starlette.websockets import WebSocketDisconnect logger = logging.getLogger("obsigate.collab") #: Délai (secondes) sans modification avant écriture sur disque. SAVE_DEBOUNCE_SECONDS = 2.0 #: Taille maximale d'une mise à jour Yjs encodée (protection anti-abus). MAX_UPDATE_BYTES = 8 * 1024 * 1024 #: Taille maximale d'un snapshot texte (protection anti-abus). MAX_TEXT_CHARS = 8 * 1024 * 1024 #: Taille maximale d'un message brut reçu (protection anti-abus, BUG-036). MAX_MESSAGE_CHARS = 16 * 1024 * 1024 #: Palette de couleurs attribuées aux utilisateurs (curseurs + avatars). PEER_COLORS = [ "#e6194b", "#3cb44b", "#4363d8", "#f58231", "#911eb4", "#008080", "#9a6324", "#800000", "#808000", "#000075", ] def color_for_index(index: int) -> str: """Return a deterministic cursor color for a peer index.""" return PEER_COLORS[index % len(PEER_COLORS)] def _b64encode(data: bytes) -> str: return base64.b64encode(data).decode("ascii") def _b64decode(data: str) -> bytes: return base64.b64decode(data.encode("ascii")) def authenticate_websocket(websocket: WebSocket) -> dict[str, Any] | None: """Authenticate a WebSocket connection. Mirrors :func:`backend.auth.middleware.get_current_user` but works on the WebSocket scope: the JWT is read from the ``access_token`` cookie, which same-origin browsers send automatically during the handshake. BUG-036: the token is **never** accepted from the query string anymore — URLs end up in access logs, proxies and browser history. Browsers cannot set custom headers on a WebSocket handshake, so the HttpOnly cookie set at login is the only supported transport. Returns the user dict, or ``None`` if authentication fails. """ from backend.auth.jwt_handler import decode_token from backend.auth.middleware import is_auth_enabled from backend.auth.user_store import get_user if not is_auth_enabled(): return { "username": "anonymous", "display_name": "Anonymous", "role": "admin", "vaults": ["*"], "active": True, "_token_vaults": ["*"], } token = websocket.cookies.get("access_token") if not token: return None payload = decode_token(token) if not payload or payload.get("type") != "access": return None user = get_user(payload["sub"]) if not user or not user.get("active"): return None user["_token_vaults"] = payload.get("vaults", []) user["_token_jti"] = payload.get("jti") return user @dataclass class CollabClient: """A single WebSocket connection inside a collaboration room.""" conn_id: int websocket: WebSocket username: str display_name: str color: str y_client_id: int | None = None awareness: dict[str, Any] | None = None def peer(self) -> dict[str, Any]: return { "connId": self.conn_id, "clientId": self.y_client_id, "username": self.username, "displayName": self.display_name, "color": self.color, } @dataclass class CollabRoom: """State shared by every client editing the same file.""" vault: str path: str file_path: Path initial_text: str = "" clients: dict[int, CollabClient] = field(default_factory=dict) #: Journal des mises à jour Yjs (binaires) depuis la création de la room. updates: list[bytes] = field(default_factory=list) has_updates: bool = False seed_sent: bool = False pending_text: str | None = None save_task: asyncio.Task | None = None lock: asyncio.Lock = field(default_factory=asyncio.Lock) @property def key(self) -> str: return f"{self.vault}::{self.path}" def peers(self) -> list[dict[str, Any]]: return [client.peer() for client in self.clients.values()] def awareness_snapshot(self) -> list[dict[str, Any]]: return [ {"clientId": c.y_client_id, "state": c.awareness} for c in self.clients.values() if c.y_client_id is not None and c.awareness is not None ] class CollabManager: """Manages collaboration rooms, broadcasting and disk persistence.""" def __init__(self, save_debounce: float = SAVE_DEBOUNCE_SECONDS) -> None: self._rooms: dict[str, CollabRoom] = {} self._save_debounce = save_debounce self._next_conn_id = 1 self._lock = asyncio.Lock() # -- introspection (used by tests / diagnostics) ------------------------ @property def room_count(self) -> int: return len(self._rooms) def room_peer_count(self, vault: str, path: str) -> int: room = self._rooms.get(f"{vault}::{path}") return len(room.clients) if room else 0 def get_room(self, vault: str, path: str) -> CollabRoom | None: return self._rooms.get(f"{vault}::{path}") # -- lifecycle ---------------------------------------------------------- async def connect( self, websocket: WebSocket, vault: str, path: str, file_path: Path, user: dict[str, Any], ) -> None: """Register *websocket* in the room and relay messages until it closes.""" async with self._lock: key = f"{vault}::{path}" room = self._rooms.get(key) if room is None: try: initial_text = file_path.read_text(encoding="utf-8") except (OSError, UnicodeDecodeError): initial_text = "" room = CollabRoom(vault=vault, path=path, file_path=file_path, initial_text=initial_text) self._rooms[key] = room conn_id = self._next_conn_id self._next_conn_id += 1 client = CollabClient( conn_id=conn_id, websocket=websocket, username=user.get("username", "anonymous"), display_name=user.get("display_name") or user.get("username", "anonymous"), color=color_for_index(conn_id - 1), ) room.clients[conn_id] = client seed: str | None = None if not room.has_updates and not room.seed_sent: seed = room.initial_text room.seed_sent = True await websocket.send_json({ "type": "init", "connId": conn_id, "color": client.color, "seed": seed, "updates": [_b64encode(u) for u in room.updates], "peers": room.peers(), "awareness": room.awareness_snapshot(), }) await self._broadcast(room, {"type": "peer_joined", "peer": client.peer()}, exclude=conn_id) try: while True: raw = await websocket.receive_text() await self._on_message(room, client, raw) except WebSocketDisconnect: pass except Exception as exc: # pragma: no cover - defensive logger.debug("Collab connection error (%s): %s", room.key, exc) finally: await self.disconnect(room, client) async def disconnect(self, room: CollabRoom, client: CollabClient) -> None: """Remove *client* from *room*, flushing and cleaning up if empty.""" async with self._lock: room.clients.pop(client.conn_id, None) empty = not room.clients if empty: await self._flush(room) async with self._lock: # Only delete if nobody rejoined while we were flushing. if not room.clients and self._rooms.get(room.key) is room: if room.save_task: room.save_task.cancel() self._rooms.pop(room.key, None) else: await self._broadcast( room, { "type": "peer_left", "peer": client.peer(), "clientId": client.y_client_id, }, ) async def stop(self) -> None: """Flush and cancel every room (called on application shutdown).""" async with self._lock: rooms = list(self._rooms.values()) self._rooms.clear() for room in rooms: if room.save_task: room.save_task.cancel() await self._flush(room) # -- message handling --------------------------------------------------- async def _on_message(self, room: CollabRoom, client: CollabClient, raw: str) -> None: # BUG-036: drop oversized frames before parsing them. if not isinstance(raw, str) or len(raw) > MAX_MESSAGE_CHARS: return try: message = json.loads(raw) except (ValueError, TypeError): return if not isinstance(message, dict): return msg_type = message.get("type") if msg_type in ("sync", "update"): encoded = message.get("update") if not isinstance(encoded, str): return try: update = _b64decode(encoded) except (ValueError, TypeError): return if not update or len(update) > MAX_UPDATE_BYTES: return async with self._lock: room.updates.append(update) room.has_updates = True await self._broadcast( room, {"type": "update", "update": encoded, "from": client.conn_id}, exclude=client.conn_id, ) elif msg_type == "awareness": y_client_id = message.get("clientId") state = message.get("state") if not isinstance(y_client_id, int): return client.y_client_id = y_client_id client.awareness = state if isinstance(state, dict) else None await self._broadcast( room, { "type": "awareness", "clientId": y_client_id, "state": client.awareness, "from": client.conn_id, }, exclude=client.conn_id, ) elif msg_type == "text": text = message.get("text") if not isinstance(text, str) or len(text) > MAX_TEXT_CHARS: return room.pending_text = text self._schedule_save(room) elif msg_type == "ping": await client.websocket.send_json({"type": "pong", "t": int(time.time() * 1000)}) # -- broadcasting ------------------------------------------------------- async def _broadcast(self, room: CollabRoom, message: dict[str, Any], exclude: int | None = None) -> None: dead: list[CollabClient] = [] for client in list(room.clients.values()): if exclude is not None and client.conn_id == exclude: continue try: await client.websocket.send_json(message) except Exception: dead.append(client) for client in dead: room.clients.pop(client.conn_id, None) # -- persistence -------------------------------------------------------- def _schedule_save(self, room: CollabRoom) -> None: if room.save_task and not room.save_task.done(): room.save_task.cancel() try: loop = asyncio.get_running_loop() except RuntimeError: # pragma: no cover - no running loop (tests) return room.save_task = loop.create_task(self._debounced_save(room)) async def _debounced_save(self, room: CollabRoom) -> None: try: await asyncio.sleep(self._save_debounce) except asyncio.CancelledError: return await self._flush(room) async def _flush(self, room: CollabRoom) -> None: """Write the last received text snapshot to disk (if any).""" async with room.lock: text = room.pending_text room.pending_text = None if text is None: return try: await asyncio.to_thread(room.file_path.write_text, text, encoding="utf-8") logger.debug("Collab persisted %s", room.key) except OSError as exc: logger.warning("Collab persist failed for %s: %s", room.key, exc) #: Process-wide singleton used by the WebSocket endpoint. collab_manager = CollabManager()