From f901e0441c762b7fabc4fc7b88c5d9d3bccc32f8 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:23:48 +0200 Subject: [PATCH] refac: bound collaborative document buffers Collaborative document buffers now keep a size budget per document and a limit on open documents per connection, only store well-formed updates and are keyed by the document id unchanged. --- backend/open_webui/socket/main.py | 44 +++++++++++++++----------- backend/open_webui/socket/utils.py | 51 ++++++++++++++++++++---------- 2 files changed, 61 insertions(+), 34 deletions(-) diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index f157ed2f81..8abca89b1b 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -666,10 +666,9 @@ async def chat_events(sid, data): def normalize_document_id(document_id: str) -> str: """Canonicalize document IDs to prevent auth bypass via prefix variants. - YdocManager normalizes storage keys by replacing ":" with "_", so - "note_abc" and "note:abc" resolve to the same underlying document. - We must rewrite underscore-prefixed IDs back to the colon form so - that authorization checks (which key on "note:") always fire. + An underscore-prefixed ID like "note_abc" would skip the authorization + checks, which key on "note:". Rewrite it back to the colon form so + those checks always fire and both forms reach the same document. """ if document_id.startswith('note_'): document_id = 'note:' + document_id[5:] @@ -710,6 +709,13 @@ async def ydoc_document_join(sid, data): user_name = data.get('user_name', 'Anonymous') user_color = data.get('user_color', '#000000') + if ( + sid not in await YDOC_MANAGER.get_users(document_id) + and await YDOC_MANAGER.count_documents_for_user(sid) >= YDOC_MANAGER.MAX_DOCUMENTS_PER_SESSION + ): + log.warning(f'Session {sid} is at the open-document limit. Rejecting join.') + return + log.info('User %s joining document %s', user_id, document_id) await YDOC_MANAGER.add_user(document_id=document_id, user_id=sid) @@ -868,23 +874,25 @@ async def yjs_document_update(sid, data): if update: user_id = data.get('user_id', sid) - await YDOC_MANAGER.append_to_updates( + stored = await YDOC_MANAGER.append_to_updates( document_id=document_id, update=update, # Convert list of bytes to bytes ) - - # Broadcast update to all other users in the document - await sio.emit( - 'ydoc:document:update', - { - 'document_id': document_id, - 'user_id': user_id, - 'update': update, - 'socket_id': sid, # Add socket_id to match frontend filtering - }, - room=f'doc_{document_id}', - skip_sid=sid, - ) + if stored: + # Broadcast update to all other users in the document + await sio.emit( + 'ydoc:document:update', + { + 'document_id': document_id, + 'user_id': user_id, + 'update': update, + 'socket_id': sid, # Add socket_id to match frontend filtering + }, + room=f'doc_{document_id}', + skip_sid=sid, + ) + else: + log.warning(f'Update for document {document_id} is invalid or over the size limit. Rejecting update.') async def debounced_save(): await asyncio.sleep(0.5) diff --git a/backend/open_webui/socket/utils.py b/backend/open_webui/socket/utils.py index f976a94e77..7cdc5490c0 100644 --- a/backend/open_webui/socket/utils.py +++ b/backend/open_webui/socket/utils.py @@ -240,6 +240,8 @@ class CachedRedisDict(RedisDict): class YdocManager: COMPACTION_THRESHOLD = 500 + MAX_DOCUMENTS_PER_SESSION = 20 + MAX_DOCUMENT_SIZE = 2 * 1024 * 1024 def __init__( self, @@ -251,20 +253,43 @@ class YdocManager: self._redis = redis self._redis_key_prefix = redis_key_prefix - async def append_to_updates(self, document_id: str, update: bytes): - document_id = document_id.replace(':', '_') + async def append_to_updates(self, document_id: str, update: list[int]) -> bool: + if not isinstance(update, list): + return False + try: + update_bytes = bytes(update) + Y.Doc().apply_update(update_bytes) # undecodable updates would stall compaction forever + except (TypeError, ValueError): + return False + update_size = len(update_bytes) + if update_size > self.MAX_DOCUMENT_SIZE: + return False if self._redis: + size_key = f'{self._redis_key_prefix}:{document_id}:size' + if await self._redis.incrby(size_key, update_size) > self.MAX_DOCUMENT_SIZE: + await self._redis.decrby(size_key, update_size) + return False redis_key = f'{self._redis_key_prefix}:{document_id}:updates' await self._redis.rpush(redis_key, JSONCodec.dumps(list(update))) list_len = await self._redis.llen(redis_key) if list_len >= self.COMPACTION_THRESHOLD: await self._compact_updates_redis(document_id) else: + if sum(len(u) for u in self._updates.get(document_id, [])) + update_size > self.MAX_DOCUMENT_SIZE: + return False if document_id not in self._updates: self._updates[document_id] = [] - self._updates[document_id].append(update) + self._updates[document_id].append(update_bytes) if len(self._updates[document_id]) >= self.COMPACTION_THRESHOLD: self._compact_updates_memory(document_id) + return True + + async def count_documents_for_user(self, user_id: str) -> int: + """Number of documents this session currently participates in.""" + if self._redis: + session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' + return await self._redis.scard(session_key) + return sum(1 for members in self._users.values() if user_id in members) async def _compact_updates_redis(self, document_id: str): """Rolling compaction: squash oldest half into one snapshot.""" @@ -276,11 +301,14 @@ class YdocManager: ydoc = Y.Doc() for raw in all_updates[:mid]: ydoc.apply_update(bytes(JSONCodec.loads(raw))) - snapshot = JSONCodec.dumps(list(ydoc.get_update())) + snapshot_list = list(ydoc.get_update()) + snapshot = JSONCodec.dumps(snapshot_list) pipe = self._redis.pipeline() pipe.delete(redis_key) pipe.rpush(redis_key, snapshot, *all_updates[mid:]) await pipe.execute() + new_size = len(snapshot_list) + sum(len(JSONCodec.loads(raw)) for raw in all_updates[mid:]) + await self._redis.set(f'{self._redis_key_prefix}:{document_id}:size', new_size) def _compact_updates_memory(self, document_id: str): """Rolling compaction: squash oldest half into one snapshot.""" @@ -294,8 +322,6 @@ class YdocManager: self._updates[document_id] = [ydoc.get_update()] + updates[mid:] async def get_updates(self, document_id: str) -> list[bytes]: - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' updates = await self._redis.lrange(redis_key, 0, -1) @@ -304,8 +330,6 @@ class YdocManager: return self._updates.get(document_id, []) async def document_exists(self, document_id: str) -> bool: - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' return await self._redis.exists(redis_key) > 0 @@ -313,8 +337,6 @@ class YdocManager: return document_id in self._updates async def get_users(self, document_id: str) -> list[str]: - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' users = await self._redis.smembers(redis_key) @@ -323,8 +345,6 @@ class YdocManager: return self._users.get(document_id, []) async def add_user(self, document_id: str, user_id: str): - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.sadd(redis_key, user_id) @@ -339,8 +359,6 @@ class YdocManager: self._users[document_id].add(user_id) async def remove_user(self, document_id: str, user_id: str): - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.srem(redis_key, user_id) @@ -350,6 +368,8 @@ class YdocManager: else: if document_id in self._users and user_id in self._users[document_id]: self._users[document_id].remove(user_id) + if not self._users[document_id]: + del self._users[document_id] async def remove_user_from_all_documents(self, user_id: str): if self._redis: @@ -380,13 +400,12 @@ class YdocManager: await self.clear_document(document_id) async def clear_document(self, document_id: str): - document_id = document_id.replace(':', '_') - if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' await self._redis.delete(redis_key) redis_users_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.delete(redis_users_key) + await self._redis.delete(f'{self._redis_key_prefix}:{document_id}:size') else: if document_id in self._updates: del self._updates[document_id]